{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-03T15:37:27.496023Z","iopub.execute_input":"2023-06-03T15:37:27.497491Z","iopub.status.idle":"2023-06-03T15:37:45.438807Z","shell.execute_reply.started":"2023-06-03T15:37:27.497451Z","shell.execute_reply":"2023-06-03T15:37:45.436974Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install pretrainedmodels\n!unzip -qq /kaggle/input/panda256/train","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:37:45.440624Z","iopub.execute_input":"2023-06-03T15:37:45.440986Z","iopub.status.idle":"2023-06-03T15:38:25.398535Z","shell.execute_reply.started":"2023-06-03T15:37:45.440955Z","shell.execute_reply":"2023-06-03T15:38:25.397094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport openslide\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\n\nimport torch\nimport cv2\nfrom tqdm import tqdm\nimport albumentations\nimport pretrainedmodels\n\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.utils.data as data_utils\nfrom torch.nn import functional as F\n\nfrom sklearn.metrics import cohen_kappa_score\nfrom fastai.losses import LabelSmoothingCrossEntropyFlat\nfrom matplotlib import pyplot as plt\n","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.401587Z","iopub.execute_input":"2023-06-03T15:38:25.402278Z","iopub.status.idle":"2023-06-03T15:38:25.41326Z","shell.execute_reply.started":"2023-06-03T15:38:25.402216Z","shell.execute_reply":"2023-06-03T15:38:25.411319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/panda256'\nDATA_DIR = '/kaggle/working/kaggle/working/train_images'\nDEVICE = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nBATCH_SIZE = 16\nEPOCHS = 30\nLEARNING_RATE = 0.02\n####\ncheckpoint_dir = '/kaggle/working/checkpoints'\nos.makedirs(checkpoint_dir, exist_ok=True)\n# Define the path to the checkpoint\ncheckpoint_path = '/kaggle/working/checkpoints/best_checkpoint.pth'","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.418196Z","iopub.execute_input":"2023-06-03T15:38:25.418702Z","iopub.status.idle":"2023-06-03T15:38:25.427916Z","shell.execute_reply.started":"2023-06-03T15:38:25.418665Z","shell.execute_reply":"2023-06-03T15:38:25.42663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(os.path.join(BASE_DIR, 'train.csv'))\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.429669Z","iopub.execute_input":"2023-06-03T15:38:25.430367Z","iopub.status.idle":"2023-06-03T15:38:25.638618Z","shell.execute_reply.started":"2023-06-03T15:38:25.430317Z","shell.execute_reply":"2023-06-03T15:38:25.636874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.640686Z","iopub.execute_input":"2023-06-03T15:38:25.641138Z","iopub.status.idle":"2023-06-03T15:38:25.650725Z","shell.execute_reply.started":"2023-06-03T15:38:25.641103Z","shell.execute_reply":"2023-06-03T15:38:25.648993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['gleason_score'].unique()","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.652667Z","iopub.execute_input":"2023-06-03T15:38:25.653254Z","iopub.status.idle":"2023-06-03T15:38:25.673641Z","shell.execute_reply.started":"2023-06-03T15:38:25.653212Z","shell.execute_reply":"2023-06-03T15:38:25.672514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['gleason_score'] = train_df['gleason_score'].str.replace('negative', '0+0')\ntrain_df['gleason_score'].unique()","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.674795Z","iopub.execute_input":"2023-06-03T15:38:25.675196Z","iopub.status.idle":"2023-06-03T15:38:25.791028Z","shell.execute_reply.started":"2023-06-03T15:38:25.675166Z","shell.execute_reply":"2023-06-03T15:38:25.789731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['gleason_score'] = train_df['gleason_score'].astype('category')\ntrain_df.dtypes","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.793149Z","iopub.execute_input":"2023-06-03T15:38:25.794466Z","iopub.status.idle":"2023-06-03T15:38:25.81685Z","shell.execute_reply.started":"2023-06-03T15:38:25.794417Z","shell.execute_reply":"2023-06-03T15:38:25.81512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mappings = dict(enumerate(train_df['gleason_score'].cat.categories))\nmappings","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.821853Z","iopub.execute_input":"2023-06-03T15:38:25.823391Z","iopub.status.idle":"2023-06-03T15:38:25.833075Z","shell.execute_reply.started":"2023-06-03T15:38:25.823338Z","shell.execute_reply":"2023-06-03T15:38:25.831637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['gleason_score'] = train_df['gleason_score'].cat.codes\n","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.834598Z","iopub.execute_input":"2023-06-03T15:38:25.835905Z","iopub.status.idle":"2023-06-03T15:38:25.849078Z","shell.execute_reply.started":"2023-06-03T15:38:25.83583Z","shell.execute_reply":"2023-06-03T15:38:25.847362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.dtypes","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.850802Z","iopub.execute_input":"2023-06-03T15:38:25.851246Z","iopub.status.idle":"2023-06-03T15:38:25.86782Z","shell.execute_reply.started":"2023-06-03T15:38:25.851211Z","shell.execute_reply":"2023-06-03T15:38:25.866451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Number of classes: {len(train_df[\"gleason_score\"].unique())}')","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.869392Z","iopub.execute_input":"2023-06-03T15:38:25.870198Z","iopub.status.idle":"2023-06-03T15:38:25.882338Z","shell.execute_reply.started":"2023-06-03T15:38:25.870155Z","shell.execute_reply":"2023-06-03T15:38:25.880724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PandaDataset(Dataset):\n    \"\"\"Custom dataset for PANDA\"\"\"\n    \n    def __init__(self, df, folds, mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)):\n        self.df = df\n        self.df = self.df[self.df.kfold.isin(folds)].reset_index(drop=True)\n        \n        # In case of validation dataset, don't apply transformations\n        if len(folds) == 1:\n            self.aug = albumentations.Compose([\n                albumentations.Normalize(mean, std, always_apply=True)\n            ])\n        else:\n            self.aug = albumentations.Compose([\n                albumentations.ShiftScaleRotate(shift_limit=0.0625,\n                                               scale_limit=0.15, \n                                               rotate_limit=10,\n                                               p=0.9),\n                albumentations.HorizontalFlip(p=0.5),\n                albumentations.VerticalFlip(p=0.5),\n                albumentations.Normalize(mean, std, always_apply=True)\n            ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        image_id = self.df.loc[index]['image_id']\n        image = cv2.imread(os.path.join(DATA_DIR, f'{image_id}.jpg'))\n        image = self.aug(image=image)['image']\n        \n        # Convert from NHWC to NCHW as pytorch expects images in NCHW format\n        image = np.transpose(image, (2, 0, 1))\n        \n        # For now, just return image and ISUP grades\n        return image, self.df.loc[index]['gleason_score']","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.884331Z","iopub.execute_input":"2023-06-03T15:38:25.884882Z","iopub.status.idle":"2023-06-03T15:38:25.900592Z","shell.execute_reply.started":"2023-06-03T15:38:25.884814Z","shell.execute_reply":"2023-06-03T15:38:25.899129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = PandaDataset(train_df, folds=[0, 1, 2, 3])\ntrain_loader = data_utils.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\n\nval_dataset = PandaDataset(train_df, folds=[4])\nval_loader = data_utils.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.90237Z","iopub.execute_input":"2023-06-03T15:38:25.903095Z","iopub.status.idle":"2023-06-03T15:38:25.943635Z","shell.execute_reply.started":"2023-06-03T15:38:25.903054Z","shell.execute_reply":"2023-06-03T15:38:25.942485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(os.path.join(BASE_DIR, 'train.csv'))\ntrain_dataset = PandaDataset(train_df, folds=[0, 1, 2, 3])\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_dataset = PandaDataset(train_df, folds=[4])\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:25.945257Z","iopub.execute_input":"2023-06-03T15:38:25.945653Z","iopub.status.idle":"2023-06-03T15:38:26.133839Z","shell.execute_reply.started":"2023-06-03T15:38:25.945611Z","shell.execute_reply":"2023-06-03T15:38:26.13279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install opencv-python\n!pip install numpy\n!pip install tqdm\n!pip install matplotlib\n!pip install easydict\n!pip install scipy\n!pip install scikit-learn","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:38:26.135314Z","iopub.execute_input":"2023-06-03T15:38:26.135892Z","iopub.status.idle":"2023-06-03T15:40:01.402148Z","shell.execute_reply.started":"2023-06-03T15:38:26.135837Z","shell.execute_reply":"2023-06-03T15:40:01.400208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:01.40512Z","iopub.execute_input":"2023-06-03T15:40:01.405633Z","iopub.status.idle":"2023-06-03T15:40:15.088146Z","shell.execute_reply.started":"2023-06-03T15:40:01.405589Z","shell.execute_reply":"2023-06-03T15:40:15.086199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision.models import resnet\nfrom timm.models.vision_transformer import vit_base_patch16_224\n\nclass ViTModel(nn.Module):\n    \"\"\"\n    Define ViT-G/14 model with 10 output classes based on Gleason scores\n    \"\"\"\n    def __init__(self, pretrained=True):\n        super(ViTModel, self).__init__()\n        if pretrained:\n            self.model = vit_base_patch16_224(pretrained=True)\n        else:\n            self.model = vit_base_patch16_224(pretrained=False)\n        \n        self.model.head = nn.Linear(self.model.head.in_features, 10)\n    \n    def forward(self, x):\n        bs, _, _, _ = x.shape\n        x = self.model(x)\n        return x\n\n\n# Usage:\npretrained = True  # Set to True if you want to use pre-trained weights\nmodel = ViTModel(pretrained)\ninput_tensor = torch.randn(1, 3, 224, 224)  # Example input tensor\noutput = model(input_tensor)\nprint(output.shape)  # Print the shape of the output tensor","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:15.090899Z","iopub.execute_input":"2023-06-03T15:40:15.091425Z","iopub.status.idle":"2023-06-03T15:40:17.850816Z","shell.execute_reply.started":"2023-06-03T15:40:15.09136Z","shell.execute_reply":"2023-06-03T15:40:17.849267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViTModel(pretrained)\nmodel.to(DEVICE) ","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:17.853196Z","iopub.execute_input":"2023-06-03T15:40:17.853643Z","iopub.status.idle":"2023-06-03T15:40:19.944002Z","shell.execute_reply.started":"2023-06-03T15:40:17.853607Z","shell.execute_reply":"2023-06-03T15:40:19.942273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass LabelSmoothingCrossEntropy(nn.Module):\n    def __init__(self, epsilon=0.1):\n        super(LabelSmoothingCrossEntropy, self).__init__()\n        self.epsilon = epsilon\n\n    def forward(self, outputs, targets):\n        num_classes = outputs.size(1)\n        device = outputs.device\n\n        one_hot = torch.zeros_like(outputs).scatter(1, targets.unsqueeze(1), 1)\n        smooth_labels = one_hot * (1 - self.epsilon) + torch.ones_like(outputs) * self.epsilon / num_classes\n\n        log_prob = torch.log_softmax(outputs, dim=1)\n\n        loss = torch.sum(-smooth_labels * log_prob, dim=1).mean()\n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:19.945718Z","iopub.execute_input":"2023-06-03T15:40:19.947108Z","iopub.status.idle":"2023-06-03T15:40:19.957932Z","shell.execute_reply.started":"2023-06-03T15:40:19.947062Z","shell.execute_reply":"2023-06-03T15:40:19.95618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.SGD(model.parameters(), lr=LEARNING_RATE)\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3, verbose=False, threshold=0.0001, threshold_mode='rel', cooldown=0, min_lr=0.00001, eps=1e-08)\ncriterion = LabelSmoothingCrossEntropy(0.1)  # Using label smoothing loss instead of normal cross entropy loss","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:19.959575Z","iopub.execute_input":"2023-06-03T15:40:19.959971Z","iopub.status.idle":"2023-06-03T15:40:19.980482Z","shell.execute_reply.started":"2023-06-03T15:40:19.959941Z","shell.execute_reply":"2023-06-03T15:40:19.978964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make use of parallel GPUs if available\nif torch.cuda.device_count() > 1:\n    model = nn.DataParallel(model)","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:19.981958Z","iopub.execute_input":"2023-06-03T15:40:19.982352Z","iopub.status.idle":"2023-06-03T15:40:19.996904Z","shell.execute_reply.started":"2023-06-03T15:40:19.982312Z","shell.execute_reply":"2023-06-03T15:40:19.994902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_accuracy(preds, targets):\n    assert len(preds) == len(targets)\n    \n    total = len(preds)\n    _, preds = torch.max(preds.data, axis=1)\n\n    correct = (preds == targets).sum().item()\n    return correct / total","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:40:19.999382Z","iopub.execute_input":"2023-06-03T15:40:20.000064Z","iopub.status.idle":"2023-06-03T15:40:20.015542Z","shell.execute_reply.started":"2023-06-03T15:40:20Z","shell.execute_reply":"2023-06-03T15:40:20.014121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torchvision.transforms import Resize\nfrom torch.nn import CrossEntropyLoss\nfrom torch.optim import SGD\nfrom torch.optim.lr_scheduler import StepLR\nfrom tqdm import tqdm\nimport os\n\n# Define the path to the checkpoint\ncheckpoint_path = '/kaggle/working/checkpoints/best_checkpoint.pth'\n# Define other variables and parameters\n\ntransform = Resize((224, 224), antialias=True)   # Resize images to (224, 224)\n\nval_acc_list = []\nbest_val_acc = 0.0  # Variable to track the best validation accuracy\n\n# Check if a checkpoint exists and load it into the model\nif os.path.exists(checkpoint_path):\n    model.load_state_dict(torch.load(checkpoint_path))\n    print(\"Checkpoint loaded successfully!\")\n\nfor epoch in range(EPOCHS):\n    model.train()\n    for i, (x_train, y_train) in tqdm(enumerate(train_loader), total=int(len(train_dataset)/train_loader.batch_size)):\n        x_train = transform(x_train)  # Resize the input images\n        x_train = x_train.to(DEVICE, dtype=torch.float32) / 255\n        y_train = y_train[0].to(DEVICE, dtype=torch.long)  # Unpack the tuple and get the labels\n        \n        # Forward pass\n        preds = model(x_train)\n        loss = criterion(preds, y_train)\n        \n        # Backpropagate\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n    lr_scheduler.step(loss.item())\n    \n    # Calculate validation accuracy after each epoch\n    # Predict on validation set\n    \n    with torch.no_grad():\n        model.eval()\n        \n        correct = 0\n        val_size = len(val_loader)\n        for x_val, y_val in tqdm(val_loader, total=int(len(val_dataset)/val_loader.batch_size)):\n            x_val = transform(x_val)  # Resize the input images\n            x_val = x_val.to(DEVICE, dtype=torch.float32) / 255\n            y_val = y_val[0].to(DEVICE, dtype=torch.long)  # Unpack the tuple and get the labels\n            \n            val_preds = model(x_val)\n            _, preds = torch.max(val_preds.data, axis=1)\n            correct += (preds == y_val).sum().item()\n        \n        val_acc = correct / len(val_dataset)\n        val_acc_list.append(val_acc)\n    \n        # Save checkpoint if the current validation accuracy is the best so far\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            checkpoint_path = os.path.join(checkpoint_dir, 'best_checkpoint.pth')\n            torch.save(model.state_dict(), checkpoint_path)\n\n    print('Epoch [{}/{}], Loss: {:.4f}, Validation accuracy: {:.2f}%'\n          .format(epoch + 1, EPOCHS, loss.item(), val_acc * 100))","metadata":{"execution":{"iopub.status.busy":"2023-06-03T15:44:04.933206Z","iopub.execute_input":"2023-06-03T15:44:04.93362Z","iopub.status.idle":"2023-06-03T15:44:05.143036Z","shell.execute_reply.started":"2023-06-03T15:44:04.933589Z","shell.execute_reply":"2023-06-03T15:44:05.140935Z"},"trusted":true},"execution_count":null,"outputs":[]}]}