{"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":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7000038,"sourceType":"datasetVersion","datasetId":4024003},{"sourceId":3729,"sourceType":"modelInstanceVersion","modelInstanceId":2656}],"dockerImageVersionId":30579,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, random_split\nimport torch.nn.functional as F\nfrom torchinfo import summary\nfrom tqdm import tqdm\nfrom sklearn.metrics import accuracy_score, balanced_accuracy_score\nfrom PIL import Image\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport os\nfrom torchvision import transforms\nimport albumentations as A\nimport albumentations.pytorch\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-19T07:38:31.710856Z","iopub.execute_input":"2023-11-19T07:38:31.711149Z","iopub.status.idle":"2023-11-19T07:38:37.549984Z","shell.execute_reply.started":"2023-11-19T07:38:31.711121Z","shell.execute_reply":"2023-11-19T07:38:37.548728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, split=\"train\", num_crops=16):\n        super().__init__()\n        self.split = split\n        self.num_crops = num_crops    # test ds \n        if split=='train':\n            self.df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\n            self.dir_path = \"/kaggle/input/UBC-OCEAN/train_thumbnails\"\n            self.alt_path = \"/kaggle/input/UBC-OCEAN/train_images\"\n        elif split==\"test\":\n            self.df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/test.csv\")\n            self.dir_path = \"/kaggle/input/UBC-OCEAN/test_thumbnails\"\n            self.alt_path = \"/kaggle/input/UBC-OCEAN/test_images\"\n        self.transformation = transforms.Compose([\n            transforms.ToTensor(),\n            transforms.Pad(32),\n            transforms.Resize((512, 512), antialias=True),\n            transforms.RandomHorizontalFlip(p=0.3),\n        ])\n        self.recrop_crit = 0.5\n        self.label_map = {'HGSC':0, 'LGSC':1, 'EC':2, 'CC':3, 'MC':4}\n        \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        if self.split == 'train':\n            image, label = self.get_data_by_id(row.image_id)\n            while True:\n                image_transformed = self.transformation(image)\n                if np.array(image_transformed == 0, dtype=np.float16).mean() < self.recrop_crit:\n                    break\n            return image_transformed, label\n        elif self.split == 'test':\n            image = self.get_data_by_id(row.image_id)\n            images = []\n            attempts = 10\n            for i in range(self.num_crops):\n                while attempts>0:\n                    image_transformed = self.transformation(image)\n                    if np.array(image_transformed == image.min(), dtype=np.float16).mean() < self.recrop_crit:\n                        break\n                    attempts -= 1\n                images.append(image_transformed)\n            images = torch.stack(images, dim=0)\n            return images, row.image_id\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def get_data_by_id(self, image_id):\n        if self.split == 'train':\n            label = self.df[self.df.image_id==image_id]['label'].item()\n            label = self.label_map[label]\n            \n        image_path = os.path.join(self.dir_path, f\"{image_id}_thumbnail.png\")\n        if os.path.exists(image_path):\n            image = cv2.imread(image_path)\n        else:\n            image_path = os.path.join(self.alt_path, f\"{image_id}.png\")\n            image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  \n        \n        image_height, image_width = image.shape[0], image.shape[1] # (width, height, 3)\n        segments = self.segments(image)\n        if len(segments) > 0:\n            i = np.random.randint(len(segments))\n            image = image[:, segments[i][0]:segments[i][1], :]\n        \n        if self.split == 'train':\n            return image, label\n        elif self.split == 'test':\n            return image\n        \n    def segments(self, image):\n        H, W = image.shape[:2]\n        vertical_sums = image.sum(axis=0).sum(axis=1).tolist()\n        vertical_sums = np.convolve(vertical_sums,np.ones(50), 'same')\n        answer = []\n        start_idx = 0 if vertical_sums[0] else None\n        for i in range(1, len(vertical_sums)):\n            if vertical_sums[i] != 0:\n                if start_idx is None:\n                    start_idx = i\n            elif start_idx is not None:\n                answer.append((start_idx, i))\n                start_idx = None\n        if start_idx is not None:\n            answer.append((start_idx, len(vertical_sums) - 1))\n        answer_remove_short = []\n        for start, end in answer:\n            if end-start > W*0.1:\n                answer_remove_short.append((start, end))\n        return answer_remove_short\n    \nds = MyDataset()\ntrain_size = int(len(ds) * 0.8)\ntrain_ds, val_ds = random_split(ds, [train_size, len(ds) - train_size])\ntest_ds = MyDataset(split='test', num_crops=1)\n# images, image_id = train_ds[10]\n# print(images.shape)\n# plt.imshow(images.cpu().detach().numpy().transpose(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:38:37.552513Z","iopub.execute_input":"2023-11-19T07:38:37.553093Z","iopub.status.idle":"2023-11-19T07:38:37.614402Z","shell.execute_reply.started":"2023-11-19T07:38:37.553052Z","shell.execute_reply":"2023-11-19T07:38:37.613668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Stratified Splitting\n# class_counts = [0] * 5\n# classMap = {'HGSC':0, 'LGSC':1, 'EC':2, 'CC':3, 'MC':4}\n# for i, row in ds.df.iterrows():\n#     label = row.label\n#     if label in class_counts:\n#         class_counts[classMap[label]] += 1\n\n# # Creating indices for each class\n# class_indices = {class_: [] for class_ in range(5)}\n# for i, row in ds.df.iterrows():\n#     label = row.label\n#     class_indices[classMap[label]].append(i)\n\n# # Split indices for each class\n# train_indices = []\n# val_indices = []\n# for _, indices in class_indices.items():\n#     random.shuffle(indices)\n#     split = int(len(indices) * 0.8)  # Example: 80% training, 20% validation\n#     train_indices += indices[:split]\n#     val_indices += indices[split:]\n\n# # Creating samplers\n# train_ds = Subset(ds, train_indices)\n# val_ds = Subset(ds, val_indices)\n\ntrain_ds, val_ds = random_split(ds, [int(len(ds)*0.8), len(ds) - int(len(ds)*0.8)])\n\ntest_ds = MyDataset(split='test', num_crops=16)\n# images, image_id = train_ds[3]\n# print(images.shape)\n# plt.imshow(images.cpu().detach().numpy().transpose(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:38:37.615343Z","iopub.execute_input":"2023-11-19T07:38:37.615582Z","iopub.status.idle":"2023-11-19T07:38:37.623348Z","shell.execute_reply.started":"2023-11-19T07:38:37.615561Z","shell.execute_reply":"2023-11-19T07:38:37.622369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\n\ntry:\n    import torch_xla\n    import torch_xla.core.xla_model as xm\n    device = xm.xla_device()\nexcept:\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nclass 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).squeeze()\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) + ')'\n    \n    \nmodel = timm.create_model('tf_efficientnet_b0', checkpoint_path='/kaggle/input/tf-efficientnet/pytorch/tf-efficientnet-b0/1/tf_efficientnet_b0_aa-827b6e33.pth')\nmodel.classifier = nn.Linear(model.classifier.in_features, 5)\nmodel.global_pool = GeM()\nmodel = model.to(device)\nsummary(model, [32, 3, 512, 512])","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:38:37.625777Z","iopub.execute_input":"2023-11-19T07:38:37.626396Z","iopub.status.idle":"2023-11-19T07:38:47.667982Z","shell.execute_reply.started":"2023-11-19T07:38:37.626362Z","shell.execute_reply":"2023-11-19T07:38:47.666901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DataLoader(dataset=train_ds, batch_size=32, shuffle=True, num_workers=4)\nval_dl = DataLoader(dataset=val_ds, batch_size=64, shuffle=False, num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:38:47.669174Z","iopub.execute_input":"2023-11-19T07:38:47.669523Z","iopub.status.idle":"2023-11-19T07:38:47.674641Z","shell.execute_reply.started":"2023-11-19T07:38:47.669496Z","shell.execute_reply":"2023-11-19T07:38:47.673585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loss_fn = nn.CrossEntropyLoss(weight = torch.FloatTensor([1/222, 1/47, 1/124, 1/99, 1/46]).to(device))\n# optimizer = torch.optim.RAdam(params=model.parameters(), lr=5e-4)\n# # scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)\n# epoch = 20\n# accumulation_step = 1\n# epoch_bal_accs = []\n# for e in range(epoch):\n#     losses = []\n#     for i, (img, label) in enumerate(train_dl):\n#         img, label = img.to(device), label.to(device)\n#         pred = model(img)\n#         loss = loss_fn(pred, label) / accumulation_step\n#         loss.backward()\n#         losses.append(loss.item() * accumulation_step)\n#         if (i + 1) % accumulation_step == 0 or i == len(train_dl) - 1:\n#             optimizer.step()\n#             optimizer.zero_grad()\n#     print(f\"Epoch {e} loss: {sum(losses) / len(losses)} \", end=\"\")\n#     with torch.no_grad():\n#         trues, preds = [], []\n#         for img, label in val_dl:\n#             img, label = img.to(device), label.to(device)\n#             pred = model(img).argmax(dim=1)\n#             trues.extend(label.tolist())\n#             preds.extend(pred.tolist())\n#         acc, bal_acc = accuracy_score(trues, preds), balanced_accuracy_score(trues, preds)\n#         epoch_bal_accs.append(bal_acc)\n#         print(f\"Accuracy: {acc} Balanced Acc: {bal_acc}\")\n#     if epoch_bal_accs[-1] == max(epoch_bal_accs):\n#         torch.save(model.state_dict(), \"/kaggle/working/model.pt\")\n# #     scheduler.step()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"/kaggle/input/efficientnet-b0-crop-1/model 4.pt\"))\nclassNames = ['HGSC', 'LGSC', 'EC', 'CC', 'MC']\nwith torch.no_grad():\n    model.eval()\n    image_ids, classes = [], []\n    for img, image_id in test_ds:\n        pred_logits = []\n        img = img.to(device)\n        pred_class = model(img).softmax(dim=1).sum(dim=0).argmax().item()\n        image_ids.append(image_id)\n        classes.append(pred_class)\n\nsubmission_df = pd.DataFrame({'image_id': image_ids, 'label': list(map(lambda x: classNames[x] , classes))})\nsubmission_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:38:47.675738Z","iopub.execute_input":"2023-11-19T07:38:47.676078Z","iopub.status.idle":"2023-11-19T07:38:50.298215Z","shell.execute_reply.started":"2023-11-19T07:38:47.676045Z","shell.execute_reply":"2023-11-19T07:38:50.29724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:39:28.88558Z","iopub.execute_input":"2023-11-19T07:39:28.886612Z","iopub.status.idle":"2023-11-19T07:39:28.895634Z","shell.execute_reply.started":"2023-11-19T07:39:28.886575Z","shell.execute_reply":"2023-11-19T07:39:28.894565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = np.random.randint(len(val_ds))\nimage, label = val_ds[i]\nplt.imshow(image.cpu().detach().numpy().transpose(1, 2, 0))\nwith torch.no_grad():\n    pred = model(image.to(device).unsqueeze(0))\nprint(pred.tolist(), label)","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:48:55.772392Z","iopub.execute_input":"2023-11-19T07:48:55.773137Z","iopub.status.idle":"2023-11-19T07:48:56.225672Z","shell.execute_reply.started":"2023-11-19T07:48:55.773105Z","shell.execute_reply":"2023-11-19T07:48:56.224792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, label = ds[315]\nds.df.iloc[315]\nimage = cv2.imread(\"/kaggle/input/UBC-OCEAN/train_thumbnails/38097_thumbnail.png\")\nimage = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\nplt.imshow(image)","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:55:28.096703Z","iopub.execute_input":"2023-11-19T07:55:28.097088Z","iopub.status.idle":"2023-11-19T07:55:29.078177Z","shell.execute_reply.started":"2023-11-19T07:55:28.097056Z","shell.execute_reply":"2023-11-19T07:55:29.077187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image = transforms.ToTensor()(image)\n# image = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])(image)\nplt.imshow(image.cpu().detach().numpy().transpose(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2023-11-19T07:55:29.07982Z","iopub.execute_input":"2023-11-19T07:55:29.080146Z","iopub.status.idle":"2023-11-19T07:55:30.107255Z","shell.execute_reply.started":"2023-11-19T07:55:29.080113Z","shell.execute_reply":"2023-11-19T07:55:30.106248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}