{"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 seaborn as sns\nimport os\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\n\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport pytorch_lightning as pl\nimport torchmetrics as tm\nimport torchvision as tv\nfrom torchvision.io import read_image\n\nfrom collections import OrderedDict","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.523906Z","iopub.execute_input":"2023-10-10T12:44:31.524266Z","iopub.status.idle":"2023-10-10T12:44:31.530114Z","shell.execute_reply.started":"2023-10-10T12:44:31.52424Z","shell.execute_reply":"2023-10-10T12:44:31.528964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DS_PATH = '/kaggle/input/UBC-OCEAN'\nSUBMISSION_MODE = True","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.537463Z","iopub.execute_input":"2023-10-10T12:44:31.538005Z","iopub.status.idle":"2023-10-10T12:44:31.5456Z","shell.execute_reply.started":"2023-10-10T12:44:31.537976Z","shell.execute_reply":"2023-10-10T12:44:31.544699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_full_df = pd.read_csv(os.path.join(DS_PATH, 'train.csv'))\ntrain_full_df.info()\ntrain_full_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.553527Z","iopub.execute_input":"2023-10-10T12:44:31.554135Z","iopub.status.idle":"2023-10-10T12:44:31.572931Z","shell.execute_reply.started":"2023-10-10T12:44:31.554106Z","shell.execute_reply":"2023-10-10T12:44:31.571912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(os.path.join(DS_PATH, 'test.csv'))\ntest_df.info()\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.574606Z","iopub.execute_input":"2023-10-10T12:44:31.575145Z","iopub.status.idle":"2023-10-10T12:44:31.591228Z","shell.execute_reply.started":"2023-10-10T12:44:31.575115Z","shell.execute_reply":"2023-10-10T12:44:31.590311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub_df = pd.read_csv(os.path.join(DS_PATH, 'sample_submission.csv'))\nsample_sub_df.info()\nsample_sub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.592763Z","iopub.execute_input":"2023-10-10T12:44:31.593319Z","iopub.status.idle":"2023-10-10T12:44:31.609056Z","shell.execute_reply.started":"2023-10-10T12:44:31.593285Z","shell.execute_reply":"2023-10-10T12:44:31.608108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EDA","metadata":{}},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_full_df.isna().sum()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.610644Z","iopub.execute_input":"2023-10-10T12:44:31.611574Z","iopub.status.idle":"2023-10-10T12:44:31.615956Z","shell.execute_reply.started":"2023-10-10T12:44:31.611541Z","shell.execute_reply":"2023-10-10T12:44:31.614945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.countplot(train_full_df, x='label')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.63324Z","iopub.execute_input":"2023-10-10T12:44:31.634146Z","iopub.status.idle":"2023-10-10T12:44:31.638752Z","shell.execute_reply.started":"2023-10-10T12:44:31.634113Z","shell.execute_reply":"2023-10-10T12:44:31.637843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.histplot(train_full_df, x='image_width', bins=20)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.658691Z","iopub.execute_input":"2023-10-10T12:44:31.658925Z","iopub.status.idle":"2023-10-10T12:44:31.663574Z","shell.execute_reply.started":"2023-10-10T12:44:31.658906Z","shell.execute_reply":"2023-10-10T12:44:31.662637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.boxplot(train_full_df, x='image_width')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.670859Z","iopub.execute_input":"2023-10-10T12:44:31.671471Z","iopub.status.idle":"2023-10-10T12:44:31.675771Z","shell.execute_reply.started":"2023-10-10T12:44:31.671443Z","shell.execute_reply":"2023-10-10T12:44:31.67486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_full_df['image_width'].describe()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.684833Z","iopub.execute_input":"2023-10-10T12:44:31.685583Z","iopub.status.idle":"2023-10-10T12:44:31.689457Z","shell.execute_reply.started":"2023-10-10T12:44:31.685554Z","shell.execute_reply":"2023-10-10T12:44:31.688702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.histplot(train_full_df, x='image_height', bins=20)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.693927Z","iopub.execute_input":"2023-10-10T12:44:31.694167Z","iopub.status.idle":"2023-10-10T12:44:31.699923Z","shell.execute_reply.started":"2023-10-10T12:44:31.694147Z","shell.execute_reply":"2023-10-10T12:44:31.699034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.boxplot(train_full_df, x='image_height')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.705169Z","iopub.execute_input":"2023-10-10T12:44:31.706035Z","iopub.status.idle":"2023-10-10T12:44:31.711031Z","shell.execute_reply.started":"2023-10-10T12:44:31.706007Z","shell.execute_reply":"2023-10-10T12:44:31.710125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_full_df['image_height'].describe()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.716931Z","iopub.execute_input":"2023-10-10T12:44:31.717206Z","iopub.status.idle":"2023-10-10T12:44:31.721988Z","shell.execute_reply.started":"2023-10-10T12:44:31.717186Z","shell.execute_reply":"2023-10-10T12:44:31.720927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.countplot(train_full_df, x='is_tma')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.728915Z","iopub.execute_input":"2023-10-10T12:44:31.729139Z","iopub.status.idle":"2023-10-10T12:44:31.733424Z","shell.execute_reply.started":"2023-10-10T12:44:31.72912Z","shell.execute_reply":"2023-10-10T12:44:31.7325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_full_df['ratio'] = train_full_df['image_width']/train_full_df['image_height']","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.737409Z","iopub.execute_input":"2023-10-10T12:44:31.737711Z","iopub.status.idle":"2023-10-10T12:44:31.743865Z","shell.execute_reply.started":"2023-10-10T12:44:31.737684Z","shell.execute_reply":"2023-10-10T12:44:31.743026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.histplot(train_full_df, x='ratio', bins=20)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.750141Z","iopub.execute_input":"2023-10-10T12:44:31.750368Z","iopub.status.idle":"2023-10-10T12:44:31.755313Z","shell.execute_reply.started":"2023-10-10T12:44:31.750348Z","shell.execute_reply":"2023-10-10T12:44:31.754282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.boxplot(train_full_df, x='ratio')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.762498Z","iopub.execute_input":"2023-10-10T12:44:31.763178Z","iopub.status.idle":"2023-10-10T12:44:31.767313Z","shell.execute_reply.started":"2023-10-10T12:44:31.763149Z","shell.execute_reply":"2023-10-10T12:44:31.766453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_full_df['ratio'].describe()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.773045Z","iopub.execute_input":"2023-10-10T12:44:31.773608Z","iopub.status.idle":"2023-10-10T12:44:31.778052Z","shell.execute_reply.started":"2023-10-10T12:44:31.77358Z","shell.execute_reply":"2023-10-10T12:44:31.777211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    sns.scatterplot(train_full_df, x='image_width', \n                    y='image_height', hue='label')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.782299Z","iopub.execute_input":"2023-10-10T12:44:31.782743Z","iopub.status.idle":"2023-10-10T12:44:31.789438Z","shell.execute_reply.started":"2023-10-10T12:44:31.782721Z","shell.execute_reply":"2023-10-10T12:44:31.78847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    nrows, ncols = 4, 4\n\n    fig, ax = plt.subplots(nrows=nrows, ncols=ncols, figsize=(10,10))\n\n    curr_img = 0\n\n    samples = train_full_df[~train_full_df['is_tma']].sample(16)\n\n    for i in range(nrows):\n        for j in range(ncols):\n            img_path = os.path.join(\n                DS_PATH, \n                'train_thumbnails',\n                str(samples.iloc[curr_img, 0])+'_thumbnail.png'\n            )\n\n            img = plt.imread(img_path)\n\n            ax[i][j].imshow(img)\n            ax[i][j].set_title(f'Image ID: {samples.iloc[curr_img, 0]},',\n                               f' Label: {samples.iloc[curr_img, 1]}')\n            ax[i][j].axis(\"off\")\n\n            curr_img += 1\n\n    plt.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.794338Z","iopub.execute_input":"2023-10-10T12:44:31.794564Z","iopub.status.idle":"2023-10-10T12:44:31.802948Z","shell.execute_reply.started":"2023-10-10T12:44:31.794545Z","shell.execute_reply":"2023-10-10T12:44:31.802097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset module","metadata":{}},{"cell_type":"code","source":"class UBCDataset(Dataset):\n    def __init__(self, \n                 annotations_df, \n                 img_dir, \n                 transform=None, \n                 target_transform=None, \n                 no_labels=False, \n                 rand_tile_size=512, \n                 file_ext='.png'):\n        self.img_labels = annotations_df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.target_transform = target_transform\n        \n        self.no_labels = no_labels\n        self.rand_tile_size = rand_tile_size\n        self.white_rgb_sum = self.rand_tile_size**2*3*255\n        \n        self.file_ext = file_ext\n\n    def __len__(self):\n        return len(self.img_labels)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(\n            self.img_dir, \n            str(self.img_labels.iloc[idx, 0])+self.file_ext\n        )\n        image = read_image(img_path)\n        \n        if image.shape[1]>self.rand_tile_size \\\n        and image.shape[2]>self.rand_tile_size:\n            is_uninformative = True\n            \n            while is_uninformative:\n                rand_h_idx = np.random.randint(\n                    image.shape[1]-self.rand_tile_size)\n\n                rand_w_idx = np.random.randint(\n                    image.shape[2]-self.rand_tile_size)\n\n                image_ = image[\n                    :,\n                    rand_h_idx:rand_h_idx+self.rand_tile_size,\n                    rand_w_idx:rand_w_idx+self.rand_tile_size\n                ]\n                \n                rgb_sum = int(image_.sum())\n                \n                if rgb_sum>self.white_rgb_sum*0.4 and rgb_sum<self.white_rgb_sum*0.9:\n                    is_uninformative = False\n                \n                #if rgb_sum>=self.white_rgb_sum*0.9:\n                #    is_uninformative = False\n                #if rgb_sum<=self.white_rgb_sum*0.4 and rgb_sum>self.white_rgb_sum*0.3:\n                #    is_uninformative = False\n            \n            image = image_\n        \n        if not self.no_labels:\n            label = torch.tensor(self.img_labels[['CC', 'EC', 'HGSC', 'LGSC', 'MC']].iloc[idx].values.astype('float'))\n        else:\n            label = []\n        \n        if self.transform:\n            image = self.transform(image)\n        if self.target_transform and not self.no_labels:\n            label = self.target_transform(label)\n            \n        return {'image':image, 'label':label}","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.816567Z","iopub.execute_input":"2023-10-10T12:44:31.817089Z","iopub.status.idle":"2023-10-10T12:44:31.826776Z","shell.execute_reply.started":"2023-10-10T12:44:31.817067Z","shell.execute_reply":"2023-10-10T12:44:31.825668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oh_label_df = pd.get_dummies(train_full_df['label'])\ntrain_full_df[oh_label_df.columns] = oh_label_df","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.848074Z","iopub.execute_input":"2023-10-10T12:44:31.848837Z","iopub.status.idle":"2023-10-10T12:44:31.856438Z","shell.execute_reply.started":"2023-10-10T12:44:31.848808Z","shell.execute_reply":"2023-10-10T12:44:31.855474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_mapping = list(oh_label_df.columns)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.858213Z","iopub.execute_input":"2023-10-10T12:44:31.858672Z","iopub.status.idle":"2023-10-10T12:44:31.863619Z","shell.execute_reply.started":"2023-10-10T12:44:31.858621Z","shell.execute_reply":"2023-10-10T12:44:31.862505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    dataset = UBCDataset(\n        train_full_df, \n        os.path.join(DS_PATH,'train_images'), \n        no_labels=False\n    )\n\n    dataloader = DataLoader(\n        dataset,\n        batch_size=1,\n        shuffle=False\n    )\n\n    data = next(iter(dataloader))\n    print(data['image'].shape)\n    print(data['label'])","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.867221Z","iopub.execute_input":"2023-10-10T12:44:31.867499Z","iopub.status.idle":"2023-10-10T12:44:31.874611Z","shell.execute_reply.started":"2023-10-10T12:44:31.867478Z","shell.execute_reply":"2023-10-10T12:44:31.873705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    tv.transforms.functional.to_pil_image(data['image'][0])","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.878931Z","iopub.execute_input":"2023-10-10T12:44:31.879463Z","iopub.status.idle":"2023-10-10T12:44:31.885122Z","shell.execute_reply.started":"2023-10-10T12:44:31.879433Z","shell.execute_reply":"2023-10-10T12:44:31.884253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#from PIL import Image\n#img = Image.open('/kaggle/input/UBC-OCEAN/train_thumbnails/4_thumbnail.png')\n#img","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.88853Z","iopub.execute_input":"2023-10-10T12:44:31.889149Z","iopub.status.idle":"2023-10-10T12:44:31.894602Z","shell.execute_reply.started":"2023-10-10T12:44:31.889129Z","shell.execute_reply":"2023-10-10T12:44:31.893696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch model training","metadata":{}},{"cell_type":"code","source":"class CustomCNN(nn.Module):\n    def __init__(self, n_classes, backbone='resnet18', \n                 requires_grad=True, pretrained_weights=True):\n        super().__init__()\n        self.backbone = backbone\n        \n        if self.backbone=='resnet18':\n            if pretrained_weights:\n                self.weights = tv.models.ResNet18_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.resnet18(weights=self.weights)\n            \n            \n        if self.backbone=='resnet34':\n            if pretrained_weights:\n                self.weights = tv.models.ResNet34_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.resnet34(weights=self.weights)\n            \n            \n        if self.backbone=='resnet50':\n            if pretrained_weights:\n                self.weights = tv.models.ResNet50_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.resnet50(weights=self.weights)\n            \n            \n        if self.backbone=='mobilenet_v3_large':\n            if pretrained_weights:\n                self.weights = tv.models.MobileNet_V3_Large_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.mobilenet_v3_large(weights=self.weights)\n            \n            \n        if self.backbone=='mobilenet_v3_small':\n            if pretrained_weights:\n                self.weights = tv.models.MobileNet_V3_Small_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.mobilenet_v3_small(weights=self.weights)\n            \n            \n        if self.backbone=='efficientnet_b0':\n            if pretrained_weights:\n                self.weights = tv.models.EfficientNet_B0_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.efficientnet_b0(weights=self.weights)\n            \n            \n        if self.backbone=='efficientnet_b1':\n            if pretrained_weights:\n                self.weights = tv.models.EfficientNet_B1_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.efficientnet_b1(weights=self.weights)\n            \n            \n        if self.backbone=='efficientnet_b2':\n            if pretrained_weights:\n                self.weights = tv.models.EfficientNet_B2_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.efficientnet_b2(weights=self.weights)\n            \n            \n        if self.backbone=='efficientnet_b3':\n            if pretrained_weights:\n                self.weights = tv.models.EfficientNet_B3_Weights.DEFAULT\n            else:\n                self.weights = None\n            self.net = tv.models.efficientnet_b3(weights=self.weights)\n        \n        \n        for param in self.net.parameters():\n            param.requires_grad = requires_grad\n        \n        if self.backbone[:6]=='resnet':\n            clf = self.net.fc\n            self.net.fc = nn.Identity()\n            self.net.fc1 = nn.Sequential(OrderedDict([\n                ('linear', nn.Linear(clf.in_features, \n                                     clf.in_features)),\n                ('relu', nn.ReLU()),\n                ('final_linear', nn.Linear(clf.in_features, n_classes))\n            ]))\n        \n        if self.backbone[:9]=='mobilenet':\n            clf = self.net.classifier\n            self.net.classifier = nn.Identity()\n            self.net.fc1 = nn.Sequential(OrderedDict([\n                ('linear', nn.Linear(clf[0].in_features, \n                                     clf[0].out_features)),\n                ('hardswish', nn.Hardswish()),\n                ('dropout', nn.Dropout(p=0.2, inplace=True)),\n                ('final_linear', nn.Linear(clf[0].out_features, n_classes))\n            ]))\n\n        if self.backbone[:12]=='efficientnet':\n            clf = self.net.classifier\n            self.net.classifier = nn.Identity()\n            self.net.fc1 = nn.Sequential(OrderedDict([\n                ('dropout', nn.Dropout(p=0.2, inplace=True)),\n                ('final_linear', nn.Linear(clf[1].in_features, n_classes))\n            ]))\n            \n    def forward(self, x):\n        out = self.net.fc1(self.net(x))\n        \n        return out\n\n\nclass UBCModel(pl.LightningModule):\n    def __init__(\n        self,\n        n_classes: int,\n        train_path,\n        valid_path,\n        test_path,\n        train_df: pd.DataFrame=None,\n        valid_df: pd.DataFrame=None,\n        test_df: pd.DataFrame=None,\n        learning_rate: float=1e-3, \n        batch_size: int=32,\n        num_workers: int=0,\n        backbone='resnet18',\n        rand_tile_size=512, \n        file_ext='.png',\n        cls_weights=None,\n        pretrained_weights=True\n    ):\n        super().__init__()\n        self.train_df = train_df\n        self.valid_df = valid_df\n        self.test_df = test_df\n\n        self.train_path = train_path\n        self.valid_path = valid_path\n        self.test_path = test_path\n        \n        self.rand_tile_size = rand_tile_size\n        self.file_ext = file_ext\n\n        self.model = CustomCNN(n_classes, backbone=backbone, \n                               pretrained_weights=pretrained_weights)\n\n        self.lr = learning_rate\n\n        self.batch_size = batch_size\n        \n        self.cls_weights = cls_weights\n\n        self.train_auroc = tm.AUROC(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        self.train_acc = tm.Accuracy(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        \n        self.val_auroc = tm.AUROC(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        self.val_acc = tm.Accuracy(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        \n        self.test_auroc = tm.AUROC(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        self.test_acc = tm.Accuracy(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n        self.test_cm = tm.ConfusionMatrix(task='multiclass', num_classes=n_classes)#, compute_on_step=False)\n\n        self.num_workers = num_workers\n\n    def forward(self, x):\n        out = self.model(x.float())\n        return out\n    \n    def configure_optimizers(self):\n        return torch.optim.Adam(\n            [p for p in self.parameters() if p.requires_grad], \n            lr=self.lr, eps=1e-08)\n    \n    def train_dataloader(self):\n        dataset = UBCDataset(self.train_df, \n                             self.train_path, \n                             rand_tile_size=self.rand_tile_size, \n                             file_ext=self.file_ext)\n        loader = DataLoader(\n            dataset,\n            batch_size=self.batch_size,\n            shuffle=True,\n            num_workers=self.num_workers\n        )\n        return loader\n\n    def val_dataloader(self):\n        dataset = UBCDataset(self.valid_df, \n                             self.valid_path, \n                             rand_tile_size=self.rand_tile_size, \n                             file_ext=self.file_ext)\n        loader = DataLoader(\n            dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=self.num_workers\n        )\n        return loader\n\n    def test_dataloader(self):\n        dataset = UBCDataset(self.test_df, \n                             self.test_path, \n                             rand_tile_size=self.rand_tile_size, \n                             file_ext=self.file_ext)\n        loader = DataLoader(\n            dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=self.num_workers\n        )\n        return loader\n    \n    def predict_dataloader(self):\n        dataset = UBCDataset(self.test_df, \n                             self.test_path, \n                             rand_tile_size=self.rand_tile_size, \n                             file_ext=self.file_ext, \n                             no_labels=True)\n        loader = DataLoader(\n            dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=self.num_workers\n        )\n        return loader\n    \n    def _common_step(self, batch):\n        output = self.forward(batch['image'])\n        \n        if self.cls_weights is not None:\n            loss_fcn = nn.CrossEntropyLoss(weight=self.cls_weights.float().to(self.device))\n        else:\n            loss_fcn = nn.CrossEntropyLoss()\n\n        loss = loss_fcn(output, batch['label'])\n        \n        return loss, output\n        \n    def training_step(self, batch, batch_idx):\n        loss, logits = self._common_step(batch)\n        \n        labels = batch['label'].argmax(dim=1)\n        \n        self.train_auroc.update(logits, labels)\n        self.train_acc.update(logits, labels)\n        \n        self.log_dict(\n            {\n                'train/loss': loss, \n                'train/auroc': self.train_auroc, \n                'train/acc': self.train_acc\n            }, \n            on_epoch=True, \n            on_step=False,\n            prog_bar=True\n        )\n        \n        return loss\n\n    def on_train_epoch_end(self):\n        self.train_acc.reset()\n        self.train_auroc.reset()\n        \n        self.val_acc.reset()\n        self.val_auroc.reset()\n    \n    def validation_step(self, batch, batch_idx):        \n        loss, logits = self._common_step(batch)\n        \n        labels = batch['label'].argmax(dim=1)\n        \n        self.val_auroc.update(logits, labels)\n        self.val_acc.update(logits, labels)\n        \n        self.log_dict(\n            {\n                'val/loss': loss, \n                'val/auroc': self.val_auroc, \n                'val/acc': self.val_acc\n            }, \n            on_epoch=True, \n            on_step=False,\n            prog_bar=True\n        )\n\n    def plot_confusion_matrix(self, df):\n        plt.figure(figsize=(4,3))\n        ax = sns.heatmap(df, annot=True, cmap='magma', fmt='')\n        ax.set_title(f'Confusion Matrix (Epoch {self.current_epoch+1})')\n        ax.set_ylabel('True labels')\n        ax.set_xlabel('Predicted labels')\n        plt.show()\n\n    def test_step(self, batch, batch_idx):        \n        loss, logits = self._common_step(batch)\n        \n        labels = batch['label'].argmax(dim=1)\n        \n        self.test_auroc.update(logits, labels)\n        self.test_acc.update(logits, labels)\n        \n        self.log_dict(\n            {\n                'test/loss': loss, \n                'test/auroc': self.test_auroc, \n                'test/acc': self.test_acc\n            }, \n            on_epoch=True, \n            on_step=False,\n            prog_bar=True\n        )\n        \n        self.test_cm.update(logits, labels)\n\n    def on_test_epoch_end(self):\n        self.plot_confusion_matrix(\n            pd.DataFrame(self.test_cm.compute().detach().cpu().numpy().astype(int)))\n\n    def predict_step(self, batch, batch_idx):\n        return self(batch['image'])\n\n\ndef train_ptl_model(\n    model,\n    model_name,\n    version,\n    epochs,\n    print_model=False\n):\n\n    if print_model:\n        print(model)\n\n    tb_logger = pl.loggers.TensorBoardLogger(save_dir='logs', \n                                             name=model_name,\n                                             version=version)\n\n    csv_logger = pl.loggers.CSVLogger(save_dir='logs', \n                                      name=model_name,\n                                      version=version)\n\n    checkpoint_callback = pl.callbacks.ModelCheckpoint(\n        dirpath=f'logs/{model_name}/{version}/best_ckpt',\n        filename=model_name+'_epoch{epoch:02d}-val_loss{val/loss:.2f}',\n        auto_insert_metric_name=False,\n        monitor='val/loss'\n    )\n\n    trainer = pl.Trainer(max_epochs=epochs, logger=[tb_logger, csv_logger], \n                         callbacks=[checkpoint_callback])\n\n    trainer.fit(model)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.900354Z","iopub.execute_input":"2023-10-10T12:44:31.900644Z","iopub.status.idle":"2023-10-10T12:44:31.933364Z","shell.execute_reply.started":"2023-10-10T12:44:31.90062Z","shell.execute_reply":"2023-10-10T12:44:31.932486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, valid_df = train_test_split(train_full_df, test_size=0.2, random_state=0)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.934845Z","iopub.execute_input":"2023-10-10T12:44:31.935613Z","iopub.status.idle":"2023-10-10T12:44:31.950766Z","shell.execute_reply.started":"2023-10-10T12:44:31.935583Z","shell.execute_reply":"2023-10-10T12:44:31.949873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def generate_class_weights(class_series, multi_class=True, one_hot_encoded=False):\n    \"\"\"\n    Method to generate class weights given a set of multi-class or multi-label labels, \n    both one-hot-encoded or not. Some examples of different formats of class_series \n    and their outputs are:\n    - generate_class_weights(['mango', 'lemon', 'banana', 'mango'], multi_class=True, \n    one_hot_encoded=False)\n    {'banana': 1.3333333333333333, 'lemon': 1.3333333333333333, \n    'mango': 0.6666666666666666}\n    - generate_class_weights([[1, 0, 0], [0, 1, 0], [0, 0, 1], [1, 0, 0]], \n    multi_class=True, one_hot_encoded=True)\n    {0: 0.6666666666666666, 1: 1.3333333333333333, 2: 1.3333333333333333}\n    - generate_class_weights([['mango', 'lemon'], ['mango'], ['lemon', 'banana'], \n    ['lemon']], multi_class=False, one_hot_encoded=False)\n    {'banana': 1.3333333333333333, 'lemon': 0.4444444444444444, \n    'mango': 0.6666666666666666}\n    - generate_class_weights([[0, 1, 1], [0, 0, 1], [1, 1, 0], [0, 1, 0]], \n    multi_class=False, one_hot_encoded=True)\n    {0: 1.3333333333333333, 1: 0.4444444444444444, 2: 0.6666666666666666}\n    The output is a dictionary in the format { class_label: class_weight }. \n    In case the input is one hot encoded, the class_label would be index of \n    appareance of the label when the dataset was processed. In multi_class \n    this is np.unique(class_series) and in multi-label \n    np.unique(np.concatenate(class_series)). Author: Angel Igareta (angel@igareta.com)\n    \"\"\"\n    if multi_class:\n        # If class is one hot encoded, transform to categorical labels to use \n        # compute_class_weight   \n        if one_hot_encoded:\n            class_series = np.argmax(class_series, axis=1)\n            \n        # Compute class weights with sklearn method\n        class_labels = np.unique(class_series)\n        class_weights = compute_class_weight(\n            class_weight='balanced', \n            classes=class_labels,\n            y=class_series\n        )\n        \n        return dict(zip(class_labels, class_weights))\n    \n    else:\n        # It is neccessary that the multi-label values are one-hot encoded\n        mlb = None\n        if not one_hot_encoded:\n            mlb = MultiLabelBinarizer()\n            class_series = mlb.fit_transform(class_series)\n\n        n_samples = len(class_series)\n        n_classes = len(class_series[0])\n\n        # Count each class frequency\n        class_count = [0] * n_classes\n        for classes in class_series:\n            for index in range(n_classes):\n                if classes[index] != 0:\n                    class_count[index] += 1\n\n        # Compute class weights using balanced method\n        class_weights = [n_samples / (n_classes * freq) \\\n                         if freq > 0 else 1 for freq in class_count]\n        class_labels = range(len(class_weights)) if mlb is None else mlb.classes_\n        \n        return dict(zip(class_labels, class_weights))","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.95223Z","iopub.execute_input":"2023-10-10T12:44:31.952826Z","iopub.status.idle":"2023-10-10T12:44:31.963267Z","shell.execute_reply.started":"2023-10-10T12:44:31.952798Z","shell.execute_reply":"2023-10-10T12:44:31.962374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cls_weights = generate_class_weights(\n    train_df[oh_label_df.columns].values, \n    multi_class=True, \n    one_hot_encoded=True\n)\n\ncls_weights = torch.tensor(list(cls_weights.values()))","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.965336Z","iopub.execute_input":"2023-10-10T12:44:31.965969Z","iopub.status.idle":"2023-10-10T12:44:31.981061Z","shell.execute_reply.started":"2023-10-10T12:44:31.965935Z","shell.execute_reply":"2023-10-10T12:44:31.97986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"model = UBCModel(\n    n_classes=train_df.label.nunique(),\n    train_path=os.path.join(DS_PATH, 'train_images'),\n    valid_path=os.path.join(DS_PATH, 'train_images'),\n    test_path=os.path.join(DS_PATH, 'test_images'),\n    train_df=train_df,\n    valid_df=valid_df,\n    test_df=test_df,\n    learning_rate=1e-3, \n    batch_size=4, \n    num_workers=0,\n    backbone='resnet18',\n    rand_tile_size=512, \n    file_ext='.png',\n    cls_weights=cls_weights,\n    pretrained_weights=True\n)\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:31.982338Z","iopub.execute_input":"2023-10-10T12:44:31.983282Z","iopub.status.idle":"2023-10-10T12:44:31.993475Z","shell.execute_reply.started":"2023-10-10T12:44:31.983251Z","shell.execute_reply":"2023-10-10T12:44:31.992391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UBCModel(\n    n_classes=train_df.label.nunique(),\n    train_path=os.path.join(DS_PATH, 'train_thumbnails'),\n    valid_path=os.path.join(DS_PATH, 'train_thumbnails'),\n    test_path=os.path.join(DS_PATH, 'test_thumbnails'),\n    train_df=train_df[~train_df['is_tma']],\n    valid_df=valid_df[~valid_df['is_tma']],\n    test_df=test_df,\n    learning_rate=1e-3, \n    batch_size=16, \n    num_workers=0,\n    backbone='resnet18',\n    rand_tile_size=512, \n    file_ext='_thumbnail.png',\n    cls_weights=cls_weights,\n    pretrained_weights=not SUBMISSION_MODE\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:47:42.254215Z","iopub.execute_input":"2023-10-10T12:47:42.254552Z","iopub.status.idle":"2023-10-10T12:47:42.428905Z","shell.execute_reply.started":"2023-10-10T12:47:42.254526Z","shell.execute_reply":"2023-10-10T12:47:42.427947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    train_ptl_model(\n        model=model,\n        model_name='ubc',\n        version='resnet18',\n        epochs=1,\n        print_model=True\n    )","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:32.206534Z","iopub.execute_input":"2023-10-10T12:44:32.206896Z","iopub.status.idle":"2023-10-10T12:44:32.211435Z","shell.execute_reply.started":"2023-10-10T12:44:32.206867Z","shell.execute_reply":"2023-10-10T12:44:32.210583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BEST_CKPT_PATH = '/kaggle/input/ubc-resnet18/ubc_epoch08-val_loss1.02.ckpt'","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:32.212682Z","iopub.execute_input":"2023-10-10T12:44:32.213604Z","iopub.status.idle":"2023-10-10T12:44:32.223728Z","shell.execute_reply.started":"2023-10-10T12:44:32.213575Z","shell.execute_reply":"2023-10-10T12:44:32.222719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"if not SUBMISSION_MODE:\n    trainer = pl.Trainer()\n\n    trainer.test(model, model.val_dataloader(), ckpt_path=BEST_CKPT_PATH)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:32.224952Z","iopub.execute_input":"2023-10-10T12:44:32.225613Z","iopub.status.idle":"2023-10-10T12:44:32.235512Z","shell.execute_reply.started":"2023-10-10T12:44:32.225583Z","shell.execute_reply":"2023-10-10T12:44:32.234696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"trainer = pl.Trainer()\n\npred = trainer.predict(model, ckpt_path=BEST_CKPT_PATH)","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:32.41568Z","iopub.execute_input":"2023-10-10T12:44:32.416197Z","iopub.status.idle":"2023-10-10T12:44:33.906411Z","shell.execute_reply.started":"2023-10-10T12:44:32.416165Z","shell.execute_reply":"2023-10-10T12:44:33.905583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_output(output_tensor: torch.tensor) -> list:\n    final_pred_list = []\n\n    for pred_batch in output_tensor:\n        batch_list = torch.nn.functional.softmax(pred_batch, dim=1)\\\n        .argmax(dim=1).tolist()\n\n        for label in batch_list:\n            final_pred_list.append(label)\n        \n    return final_pred_list","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:33.907851Z","iopub.execute_input":"2023-10-10T12:44:33.90818Z","iopub.status.idle":"2023-10-10T12:44:33.913693Z","shell.execute_reply.started":"2023-10-10T12:44:33.90815Z","shell.execute_reply":"2023-10-10T12:44:33.912508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['label'] = process_output(pred)\ntest_df['label'] = test_df['label'].apply(lambda x: class_mapping[x])\nsubmission_df = test_df.drop(columns=['image_width', 'image_height'])\nsubmission_df.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:33.914934Z","iopub.execute_input":"2023-10-10T12:44:33.915804Z","iopub.status.idle":"2023-10-10T12:44:33.928463Z","shell.execute_reply.started":"2023-10-10T12:44:33.915774Z","shell.execute_reply":"2023-10-10T12:44:33.927674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-10T12:44:33.929629Z","iopub.execute_input":"2023-10-10T12:44:33.929968Z","iopub.status.idle":"2023-10-10T12:44:33.946881Z","shell.execute_reply.started":"2023-10-10T12:44:33.929938Z","shell.execute_reply":"2023-10-10T12:44:33.945795Z"},"trusted":true},"execution_count":null,"outputs":[]}]}