{"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":"markdown","source":"---\n<h3 align=\"center\"> <strong><h1 align=\"center\"><strong>Ovarian Cancer Subtype Classification Using ViT and PyTorch</strong></h3>\n\n---","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport PIL\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport plotly.express as px\nimport matplotlib.pyplot as plt\n\nimport datasets\nfrom datasets import load_dataset\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.transforms import (CenterCrop,\n                                    Compose,\n                                    Normalize,\n                                    RandomHorizontalFlip,\n                                    RandomResizedCrop,\n                                    Resize,\n                                    ToTensor)\n\nfrom transformers import (ViTImageProcessor,\n                          ViTForImageClassification,\n                          get_linear_schedule_with_warmup,\n                          AdamW)\n\n\nfrom sklearn import metrics\nfrom sklearn.model_selection import train_test_split\n\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2023-11-09T12:10:36.469528Z","iopub.execute_input":"2023-11-09T12:10:36.469791Z","iopub.status.idle":"2023-11-09T12:10:51.183647Z","shell.execute_reply.started":"2023-11-09T12:10:36.469766Z","shell.execute_reply":"2023-11-09T12:10:51.182867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"TRAIN_THUMBNAILS = '/kaggle/input/UBC-OCEAN/train_thumbnails'\nTRAIN_IMAGES = '/kaggle/input/UBC-OCEAN/train_images'\n\ndef get_file_path(image_id, thumbnails= TRAIN_THUMBNAILS, images=TRAIN_IMAGES):\n    if os.path.exists(f\"{thumbnails}/{image_id}_thumbnail.png\"):\n        return f\"{thumbnails}/{image_id}_thumbnail.png\"\n    else:\n        return f\"{images}/{image_id}.png\"\n    \n\ndef get_cifar_datasets():\n  train_ds, test_ds = load_dataset('cifar10', split=['train[:50]', 'test[:20]'])\n  splits = train_ds.train_test_split(test_size=0.1)\n  train_ds = splits['train']\n  valid_ds = splits['test']\n  return train_ds, valid_ds, test_ds\n\n\ndef add_images(train_df):\n    img_ids = []\n    for img_id in  tqdm(train_df.image_id, total=len(train_df)):\n#         try:\n        PIL.Image.MAX_IMAGE_PIXELS = None\n        img_ids.append(\n            PIL.Image.open(get_file_path(img_id)).resize(\n                (300,200), \n                PIL.Image.ANTIALIAS\n            ))\n#         except:\n#             img_ids.append(None)\n\n    train_df[\"image\"] = img_ids\n    train_df = train_df[train_df.image.notna()].reset_index(drop=True)\n    train_df = train_df[[\"image_id\", \"image\", \"label\"]]\n    train_df[\"label_name\"] = train_df[\"label\"]\n    return train_df\n\n\ndef plot_subtypes_dists(train_df):\n    generated_counts = train_df['label_name'].value_counts().reset_index()\n    generated_counts.columns = ['label', 'count']\n    fig = px.pie(\n        generated_counts, \n        names='label', \n        values='count', \n        title='Distribution of Ovarian Cancer Subtypes')\n    fig.show()\n\n\ndef encode_labels(train_df):\n    id2label = {id:label for id, label in enumerate(train_df['label'].unique())}\n    label2id = {label:id for id,label in id2label.items()}\n    train_df[\"label\"] = [label2id[label] for label in train_df[\"label\"]]\n    return train_df, id2label, label2id\n\n# train_ds, valid_ds, test_ds = get_cifar_datasets()","metadata":{"execution":{"iopub.status.busy":"2023-11-09T12:10:51.185167Z","iopub.execute_input":"2023-11-09T12:10:51.185498Z","iopub.status.idle":"2023-11-09T12:10:51.198088Z","shell.execute_reply.started":"2023-11-09T12:10:51.18547Z","shell.execute_reply":"2023-11-09T12:10:51.197193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\ntrain_df = add_images(train_df)\n\ntrain_df, id2label, label2id = encode_labels(train_df)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:50:17.327099Z","iopub.execute_input":"2023-11-09T11:50:17.327839Z","iopub.status.idle":"2023-11-09T11:53:09.655555Z","shell.execute_reply.started":"2023-11-09T11:50:17.327805Z","shell.execute_reply":"2023-11-09T11:53:09.65462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_subtypes_dists(train_df)","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:53:49.968386Z","iopub.execute_input":"2023-11-09T11:53:49.968794Z","iopub.status.idle":"2023-11-09T11:53:50.024112Z","shell.execute_reply.started":"2023-11-09T11:53:49.968765Z","shell.execute_reply":"2023-11-09T11:53:50.023263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class Config:\n    EPOCHS = 4\n    TRAIN_BATCH_SIZE = 8\n    VALID_BATCH_SIZE = 4\n    LEARNING_RATE = 5e-5\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    MODEL_CHECKPOINT = \"/kaggle/input/hugging-face-google-vit/google-vit-base-patch16-224-in21k\"\n    \n#     MODEL_PATH = None\n#     TRAIN_DATA = None\n\n    PROCESSOR = ViTImageProcessor.from_pretrained(MODEL_CHECKPOINT)\n    VIT_MODEL = ViTForImageClassification.from_pretrained(MODEL_CHECKPOINT,\n                                                          num_labels=len(id2label),\n                                                          id2label=id2label,\n                                                          label2id=label2id)\n\n    IMG_MEAN = PROCESSOR.image_mean\n    IMG_STD = PROCESSOR.image_std\n    SIZE = PROCESSOR.size[\"height\"]","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:19.917267Z","iopub.execute_input":"2023-11-09T11:54:19.917664Z","iopub.status.idle":"2023-11-09T11:54:21.034776Z","shell.execute_reply.started":"2023-11-09T11:54:19.917636Z","shell.execute_reply":"2023-11-09T11:54:21.033954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_mean = Config.IMG_MEAN\nimage_std = Config.IMG_STD\nsize = Config.SIZE\n\n\ndef transforms(image, image_mean=image_mean, image_std=image_std, size=size):\n    normalize = Normalize(mean=image_mean, std=image_std)\n    _transforms = Compose([\n        RandomResizedCrop(size),\n        RandomHorizontalFlip(),\n        ToTensor(),\n        normalize])\n    return _transforms(image.convert(\"RGB\"))","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:24.887284Z","iopub.execute_input":"2023-11-09T11:54:24.88774Z","iopub.status.idle":"2023-11-09T11:54:24.894657Z","shell.execute_reply.started":"2023-11-09T11:54:24.887703Z","shell.execute_reply":"2023-11-09T11:54:24.893548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class OvarianDataset(Dataset):\n    def __init__(self, img_paths, img_labels, transforms=None):\n        self.img_paths = img_paths\n        self.img_labels = img_labels\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def __getitem__(self, index):\n        img_path = self.img_paths[index]\n        img_label = self.img_labels[index]\n\n        # image = PIL.Image.open(img_path)\n\n        if transforms is not None:\n            pixel_values = transforms(img_path)\n\n\n        return {\n            \"pixel_values\": pixel_values,\n            \"labels\": img_label\n        }","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:27.516339Z","iopub.execute_input":"2023-11-09T11:54:27.516719Z","iopub.status.idle":"2023-11-09T11:54:27.523794Z","shell.execute_reply.started":"2023-11-09T11:54:27.516689Z","shell.execute_reply":"2023-11-09T11:54:27.522772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Engine","metadata":{}},{"cell_type":"code","source":"def loss_fn(logits, labels):\n    criterion = nn.CrossEntropyLoss()\n    loss = criterion(logits, labels)\n    return loss\n\n\ndef train_fn(train_dataloader, model, optimizer, scheduler, device):\n    model.train()\n    final_loss = 0\n\n    fin_labels = []\n    fin_outputs = []\n\n    for data in tqdm(train_dataloader, total=len(train_dataloader)):\n        pixel_values = data[\"pixel_values\"]\n        labels = data[\"labels\"]\n\n        pixel_values = pixel_values.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n        logits, loss = model(pixel_values, labels)\n        final_loss += loss.item()\n\n        fin_labels.extend([id2label[id.item()] for id in labels.cpu().detach()])\n        fin_outputs.extend([id2label[id.item()] for id in logits.argmax(-1).cpu().detach()])\n\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n\n    return fin_labels, fin_outputs, final_loss/len(train_dataloader)\n\n\n\ndef valid_fn(valid_dataloader, model, device):\n    model.eval()\n    final_loss = 0\n\n    fin_labels = []\n    fin_outputs = []\n    with torch.no_grad():\n        for data in tqdm(valid_dataloader, total=len(valid_dataloader)):\n            pixel_values = data[\"pixel_values\"]\n            labels = data[\"labels\"]\n\n            pixel_values = pixel_values.to(device)\n            labels = labels.to(device)\n\n            logits, loss = model(pixel_values, labels)\n            final_loss += loss.item()\n\n            fin_labels.extend([id2label[id.item()] for id in labels.cpu().detach()])\n            fin_outputs.extend([id2label[id.item()] for id in logits.argmax(-1).cpu().detach()])\n\n    return fin_outputs, fin_labels, final_loss/len(valid_dataloader)","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:28.960485Z","iopub.execute_input":"2023-11-09T11:54:28.96133Z","iopub.status.idle":"2023-11-09T11:54:28.973097Z","shell.execute_reply.started":"2023-11-09T11:54:28.961299Z","shell.execute_reply":"2023-11-09T11:54:28.972138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class ViTOvarianModel(nn.Module):\n    def __init__(self):\n        super(ViTOvarianModel, self).__init__()\n        self.vit = Config.VIT_MODEL\n\n    def forward(self, pixel_values, labels=None):\n        logits = self.vit(pixel_values=pixel_values).logits\n\n        loss = None\n        if labels is not None:\n            loss = loss_fn(logits, labels)\n\n        return logits, loss","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:30.268256Z","iopub.execute_input":"2023-11-09T11:54:30.268633Z","iopub.status.idle":"2023-11-09T11:54:30.274622Z","shell.execute_reply.started":"2023-11-09T11:54:30.268592Z","shell.execute_reply":"2023-11-09T11:54:30.273667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"train_ds, valid_ds = train_test_split(train_df, test_size=0.1, random_state=42)\n\ntrain_ds.reset_index(drop=True, inplace=True)\nvalid_ds.reset_index(drop=True, inplace=True)\n\ntrain_dataset = OvarianDataset(\n    train_ds[\"image\"],\n    train_ds[\"label\"],\n    )\n\nvalid_dataset = OvarianDataset(\n    valid_ds[\"image\"],\n    valid_ds[\"label\"],\n    )\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=Config.TRAIN_BATCH_SIZE, num_workers=4)\nvalid_dataloader = DataLoader(valid_dataset, batch_size=Config.VALID_BATCH_SIZE, num_workers=1)\n\n\nmodel = ViTOvarianModel()\nmodel.to(Config.DEVICE)\n\nnum_train_steps = int(len(train_dataset) / Config.TRAIN_BATCH_SIZE * Config.EPOCHS)\noptimizer  = AdamW(model.parameters(), lr=Config.LEARNING_RATE)\nscheduler = get_linear_schedule_with_warmup(\n    optimizer, num_warmup_steps=0, num_training_steps=num_train_steps\n    )","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:32.649872Z","iopub.execute_input":"2023-11-09T11:54:32.650235Z","iopub.status.idle":"2023-11-09T11:54:37.886416Z","shell.execute_reply.started":"2023-11-09T11:54:32.650208Z","shell.execute_reply":"2023-11-09T11:54:37.885362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = defaultdict(list)\nbest_accuracy = 0\n\nfor epoch in range(1, Config.EPOCHS+1):\n    train_outputs, train_labels, train_loss = train_fn(train_dataloader, model, optimizer, scheduler, Config.DEVICE)\n    valid_outputs, valid_labels, valid_loss = valid_fn(valid_dataloader, model, Config.DEVICE)\n\n    train_accuracy = metrics.accuracy_score(train_labels, train_outputs)\n    valid_accuracy = metrics.accuracy_score(valid_labels, valid_outputs)\n\n    print(f\"Epoch: {epoch}\\nTrain Loss: {train_loss} - Train Accuracy: {train_accuracy} \\nValid Loss: {valid_loss} - Valid Accuracy: {valid_accuracy}\\n\")\n\n    history['Train Loss'].append(train_loss)\n    history['Train Accuracy'].append(train_accuracy)\n    history['Valid Loss'].append(valid_loss)\n    history['Valid Accuracy'].append(valid_accuracy)\n\n#     if valid_accuracy > best_accuracy:\n#         torch.save(model.state_dict(), Config.MODEL_PATH)\n#         best_accuracy = valid_accuracy","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:54:37.888369Z","iopub.execute_input":"2023-11-09T11:54:37.888726Z","iopub.status.idle":"2023-11-09T11:55:25.982069Z","shell.execute_reply.started":"2023-11-09T11:54:37.888694Z","shell.execute_reply":"2023-11-09T11:55:25.980988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"def get_metrics(labels, outputs, avg = \"macro\"):\n  accuracy = metrics.accuracy_score(labels, outputs)\n  recall = metrics.recall_score(labels, outputs, average=avg)\n  precision = metrics.precision_score(labels, outputs, average=avg)\n  f1 = metrics.f1_score(labels, outputs, average=avg)\n  return accuracy, recall, precision, f1","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:55:29.956325Z","iopub.execute_input":"2023-11-09T11:55:29.956715Z","iopub.status.idle":"2023-11-09T11:55:29.962509Z","shell.execute_reply.started":"2023-11-09T11:55:29.956685Z","shell.execute_reply":"2023-11-09T11:55:29.961468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accuracy, recall, precision, f1 = get_metrics(valid_labels, valid_outputs)\n\nprint(\"===\"*50)\nprint(\"\\nResults summary\\n\")\nprint(f\"Accuracy Score  : {accuracy}\")\nprint(f\"Recall Score    : {recall}\")\nprint(f\"Precision Score : {precision}\")\nprint(f\"F1 Score        : {f1}\")\n\nprint(\"===\"*50)\nprint(\"\\nClassification report \\n\\n\", metrics.classification_report(valid_labels, valid_outputs))\n\nprint(\"===\"*50)\ncm = metrics.confusion_matrix(valid_labels, valid_outputs)\nfig, ax = plt.subplots()\nsns.heatmap(cm, annot=True, fmt='d', ax=ax, cmap=plt.cm.Blues, cbar=False)\nax.set(xlabel=\"Predicted Label\",\n       ylabel=\"True Label\",\n       xticklabels=np.unique(valid_labels),\n       yticklabels=np.unique(valid_labels),\n       title=\"CONFUSION MATRIX\")\nplt.yticks(rotation=0)\nplt.xticks(rotation=45);","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:55:30.472505Z","iopub.execute_input":"2023-11-09T11:55:30.472883Z","iopub.status.idle":"2023-11-09T11:55:30.774197Z","shell.execute_reply.started":"2023-11-09T11:55:30.472854Z","shell.execute_reply":"2023-11-09T11:55:30.773252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a 1x2 grid of subplots\nfig, axs = plt.subplots(1, 2, figsize=(12, 4))\n\n# Plot the first set of data (accuracy)\naxs[0].plot(history['Train Accuracy'], '-o', label='Train Accuracy')\naxs[0].plot(history['Valid Accuracy'], '-o', label='Validation Accuracy')\naxs[0].set_title('Accuracy')\naxs[0].set_ylabel('Accuracy')\naxs[0].set_xlabel('Epoch')\naxs[0].legend()\naxs[0].set_ylim([0, 1])\n\n# Plot the second set of data (loss)\naxs[1].plot(history['Train Loss'], '-o', label='Train Loss')\naxs[1].plot(history['Valid Loss'], '-o', label='Validation Loss')\naxs[1].set_title('Loss')\naxs[1].set_ylabel('Loss')\naxs[1].set_xlabel('Epoch')\naxs[1].legend()\naxs[1].set_ylim([0, 3])\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:55:35.42915Z","iopub.execute_input":"2023-11-09T11:55:35.429822Z","iopub.status.idle":"2023-11-09T11:55:36.015746Z","shell.execute_reply.started":"2023-11-09T11:55:35.42979Z","shell.execute_reply":"2023-11-09T11:55:36.01485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"def predict(img_path, model=model, device=Config.DEVICE):\n    image = PIL.Image.open(img_path).resize((600,500), PIL.Image.ANTIALIAS)\n    pixel_values = transforms(image).unsqueeze(0)\n    pixel_values = pixel_values.to(device)\n    logits, _ = model(pixel_values)\n    pred = logits.argmax(-1).cpu().detach()\n    pred = id2label[pred.item()]\n    return pred","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:55:40.599002Z","iopub.execute_input":"2023-11-09T11:55:40.599363Z","iopub.status.idle":"2023-11-09T11:55:40.606573Z","shell.execute_reply.started":"2023-11-09T11:55:40.599336Z","shell.execute_reply":"2023-11-09T11:55:40.605584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_THUMBNAILS = '/kaggle/input/UBC-OCEAN/test_thumbnails'\nTEST_IMAGES = '/kaggle/input/UBC-OCEAN/test_images'","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:57:59.307932Z","iopub.execute_input":"2023-11-09T11:57:59.308808Z","iopub.status.idle":"2023-11-09T11:57:59.313033Z","shell.execute_reply.started":"2023-11-09T11:57:59.308774Z","shell.execute_reply":"2023-11-09T11:57:59.312109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/test.csv\")\n\ntest_df['file_path'] = test_df['image_id'].apply(get_file_path, thumbnails= TEST_THUMBNAILS, images=TEST_IMAGES)\ntest_df['label'] = test_df['file_path'].apply(predict)\n\ntest_df[['image_id', 'label']].to_csv('submission.csv', index = False)\ntest_df[['image_id', 'label']].head()","metadata":{"execution":{"iopub.status.busy":"2023-11-09T11:59:15.995977Z","iopub.execute_input":"2023-11-09T11:59:15.99633Z","iopub.status.idle":"2023-11-09T11:59:16.346294Z","shell.execute_reply.started":"2023-11-09T11:59:15.996304Z","shell.execute_reply":"2023-11-09T11:59:16.34533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}