{"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":7108226,"sourceType":"datasetVersion","datasetId":4098242}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport pytorch_lightning as pl\nimport os\nimport random\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import transforms\nimport torch.optim as optim\nimport torch.nn as nn\nimport torchvision\nfrom PIL import Image\nfrom torch.utils.data import DataLoader, Dataset\nfrom transformers import ViTImageProcessor, ViTForImageClassification\nimport torch.nn as nn\nimport torch.optim as optim\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-02T21:03:46.265284Z","iopub.execute_input":"2023-12-02T21:03:46.265671Z","iopub.status.idle":"2023-12-02T21:04:03.185002Z","shell.execute_reply.started":"2023-12-02T21:03:46.265641Z","shell.execute_reply":"2023-12-02T21:04:03.184225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViT(pl.LightningModule):\n    def __init__(self, num_classes, image_size=224, backbone='vit_base_patch16_224', finetune_layer=True):\n        super().__init__()\n        self.model = timm.create_model(backbone, pretrained=True)\n        \n        if finetune_layer:\n            for param in self.model.parameters():\n                param.requires_grad = False\n            self.model.head = nn.Identity()\n            self.finetune_layer = nn.Sequential(\n                nn.Linear(self.model(torch.randn(1, 3, image_size, image_size)).shape[-1], 512),\n                nn.ReLU(),\n                nn.Linear(512, num_classes)\n            )\n        else:\n            self.model.head = nn.Linear(self.model.head.in_features, num_classes)\n            self.finetune_layer = None\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        \n        if self.finetune_layer is not None:\n            logits = self.finetune_layer(logits)\n\n        loss = F.cross_entropy(logits, y)\n        self.log('train_loss', loss)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n\n        if self.finetune_layer is not None:\n            logits = self.finetune_layer(logits)\n\n        loss = F.cross_entropy(logits, y)\n        self.log('val_loss', loss)\n        return loss\n\n    def configure_optimizers(self):\n        if self.finetune_layer is not None:\n            parameters = list(self.model.parameters()) + list(self.finetune_layer.parameters())\n        else:\n            parameters = self.parameters()\n\n        optimizer = torch.optim.Adam(parameters, lr=1e-4)\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:04:16.754454Z","iopub.execute_input":"2023-12-02T21:04:16.755611Z","iopub.status.idle":"2023-12-02T21:04:16.767294Z","shell.execute_reply.started":"2023-12-02T21:04:16.755579Z","shell.execute_reply":"2023-12-02T21:04:16.766205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViT(num_classes=5)\n\ncheckpoint_path = '/kaggle/input/vit-vision-transformer-vit-base-patch16-224/vit-epoch13-val_loss1.2933.ckpt'\ncheckpoint = torch.load(checkpoint_path)\n\nmodel.load_state_dict(checkpoint['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:04:19.493557Z","iopub.execute_input":"2023-12-02T21:04:19.493912Z","iopub.status.idle":"2023-12-02T21:04:48.912047Z","shell.execute_reply.started":"2023-12-02T21:04:19.493883Z","shell.execute_reply":"2023-12-02T21:04:48.911132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_tiles(img, tile_size=256, n_tiles=30, mode=0):\n    h, w, c = img.shape\n    pad_h = (tile_size - h % tile_size) % tile_size + ((tile_size * mode) // 2)\n    pad_w = (tile_size - w % tile_size) % tile_size + ((tile_size * mode) // 2)\n\n    img = np.pad(\n        img,\n        [[pad_h // 2, pad_h - pad_h // 2], [pad_w // 2, pad_w - pad_w // 2], [0, 0]],\n        constant_values=0,\n    )\n    img = img.reshape(\n        img.shape[0] // tile_size, tile_size, img.shape[1] // tile_size, tile_size, 3\n    )\n    img = img.transpose(0, 2, 1, 3, 4).reshape(-1, tile_size, tile_size, 3)\n    \n    idxs = np.argsort(img.reshape(img.shape[0], -1).sum(-1))\n    if len(img) < n_tiles:\n        img = np.pad(\n            img, [[0, n_tiles - len(img)], [0,0], [0,0], [0,0]], constant_values=255\n        )\n    # idxs = np.argsort(-img.reshape(img.shape[0], -1).sum(-1))[:n_tiles]\n    # print(type(idxs))\n    if idxs.shape[0]>n_tiles:\n        idxs = idxs[-n_tiles:]\n    img = img[idxs]\n    \n    return img\n\ndef concat_tiles(tiles, n_tiles, image_size):\n    idxes = list(range(n_tiles))\n    \n    n_row_tiles = int(np.sqrt(n_tiles))\n    img = np.zeros(\n        (image_size*n_row_tiles, image_size*n_row_tiles, 3), dtype=\"uint8\"\n    )\n    \n    for h in range(n_row_tiles):\n        for w in range(n_row_tiles):\n            i = h * n_row_tiles + w\n            if len(tiles) > idxes[i]:\n                this_img = tiles[idxes[i]]\n            else:\n                this_img = np.ones((image_size, image_size, 3), dtype=\"uint8\") * 255\n                \n            h1 = h * image_size\n            w1 = w * image_size\n            img[h1 : h1 + image_size, w1 : w1 + image_size] = this_img\n    return img\n\ndef sort_tiles_by_intensity(tiles):\n    intensities = np.mean(tiles, axis=(1, 2, 3))  # Calculate mean intensity for each tile\n    \n    sorted_indices = np.argsort(-intensities)\n    \n    sorted_tiles = tiles[sorted_indices]\n    sorted_intensities = intensities[sorted_indices]\n    \n    return sorted_tiles, sorted_intensities, sorted_indices\n","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:04:55.697148Z","iopub.execute_input":"2023-12-02T21:04:55.697986Z","iopub.status.idle":"2023-12-02T21:04:55.711398Z","shell.execute_reply.started":"2023-12-02T21:04:55.697954Z","shell.execute_reply":"2023-12-02T21:04:55.710478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UCBDataset(Dataset):\n    def __init__(self, metadata_df, image_folder, transform=None):\n        self.metadata_df = metadata_df\n        self.image_folder = image_folder\n        self.transform = transform  # Use the provided transform\n    def __len__(self):\n        return len(self.metadata_df)\n    def __getitem__(self, idx):\n        image_ids = self.metadata_df.image_id[idx]  \n        image_name = os.path.join(self.image_folder, \"{}_thumbnail.png\".format(image_ids))\n        image = Image.open(image_name)\n        img = get_tiles(\n            np.array(image),\n            mode=0, n_tiles= 64\n        )\n        \n        sorted_img = sort_tiles_by_intensity(img)\n        img = concat_tiles(\n            img, 64, 256\n        )\n\n        # img = to_tensor(img)\n        img = Image.fromarray(img)\n        \n        if self.transform:\n            img = self.transform(img)\n\n        return img","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:04:58.863086Z","iopub.execute_input":"2023-12-02T21:04:58.863724Z","iopub.status.idle":"2023-12-02T21:04:58.871283Z","shell.execute_reply.started":"2023-12-02T21:04:58.863693Z","shell.execute_reply":"2023-12-02T21:04:58.870404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = transforms.Compose([\n        \n        transforms.Resize(224),\n        transforms.ToTensor(),\n        transforms.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225]\n            )\n    ])\n\ntest_df = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\ntest_dataset = UCBDataset(metadata_df=test_df, image_folder='/kaggle/input/UBC-OCEAN/test_thumbnails', transform=data_transforms)\n\nbatch_size = 4\ntest_loader = DataLoader(test_dataset, batch_size=batch_size)","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:05:03.980529Z","iopub.execute_input":"2023-12-02T21:05:03.981222Z","iopub.status.idle":"2023-12-02T21:05:03.996591Z","shell.execute_reply.started":"2023-12-02T21:05:03.981191Z","shell.execute_reply":"2023-12-02T21:05:03.995561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nall_preds = []\nmodel.to(device)\nwith torch.no_grad():\n    for inputs in tqdm(test_loader, desc='testing', leave=False):\n        inputs = inputs.to(device)\n        outputs = model(inputs)\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n    print('Testing done!')","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:05:07.183578Z","iopub.execute_input":"2023-12-02T21:05:07.183931Z","iopub.status.idle":"2023-12-02T21:05:11.830368Z","shell.execute_reply.started":"2023-12-02T21:05:07.183897Z","shell.execute_reply":"2023-12-02T21:05:11.829329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_mapping = {0: 'CC', 1: 'EC', 2: 'HGSC', 3: 'LGSC', 4: 'MC'}  # Adjust this based on your column order\n\ndecoded_preds = [label_mapping[pred] for pred in all_preds]","metadata":{"execution":{"iopub.status.busy":"2023-12-02T21:05:13.684672Z","iopub.execute_input":"2023-12-02T21:05:13.685327Z","iopub.status.idle":"2023-12-02T21:05:14.964426Z","shell.execute_reply.started":"2023-12-02T21:05:13.68529Z","shell.execute_reply":"2023-12-02T21:05:14.962972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}