{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4619805,"sourceType":"datasetVersion","datasetId":2688675},{"sourceId":4824226,"sourceType":"datasetVersion","datasetId":2794616}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pydicom\nimport pandas as pd\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom ignite.engine import Events, create_supervised_trainer, create_supervised_evaluator\nfrom ignite.metrics import Accuracy, Loss, RunningAverage\nfrom ignite.contrib.handlers import ProgressBar\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import models, transforms","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-20T16:13:08.626553Z","iopub.execute_input":"2024-05-20T16:13:08.626943Z","iopub.status.idle":"2024-05-20T16:13:15.45742Z","shell.execute_reply.started":"2024-05-20T16:13:08.626913Z","shell.execute_reply":"2024-05-20T16:13:15.456008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_xray(file_path, img_size=None):\n    \"\"\"\n    Read the dicom data and get the image\n    Args:\n        file_path: The path of the dicom file\n        img_size: Size of the output image\n    \"\"\"\n\n    dicom = pydicom.read_file(file_path)\n    img = dicom.pixel_array\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = np.max(img) - img\n\n    if img_size:\n        img = cv2.resize(img, img_size)\n\n    # Add channel dim at First\n    img = img[np.newaxis]\n\n    # Converting img to float32\n    img = img / np.max(img)\n    img = img.astype(\"float32\")\n\n    return img","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:13:15.459573Z","iopub.execute_input":"2024-05-20T16:13:15.460065Z","iopub.status.idle":"2024-05-20T16:13:15.467187Z","shell.execute_reply.started":"2024-05-20T16:13:15.460034Z","shell.execute_reply":"2024-05-20T16:13:15.466107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def patchify(batch, patch_size):\n    \"\"\"\n    Patchify the batch of images\n        \n    Shape:\n        batch: (b, h, w, c)\n        output: (b, nh, nw, ph, pw, c)\n    \"\"\"\n    b, c, h, w = batch.shape\n    ph, pw = patch_size\n    nh, nw = h // ph, w // pw\n\n    batch_patches = torch.reshape(batch, (b, c, nh, ph, nw, pw))\n    batch_patches = torch.permute(batch_patches, (0, 1, 2, 4, 3, 5))\n\n    return batch_patches","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:13:15.468569Z","iopub.execute_input":"2024-05-20T16:13:15.468868Z","iopub.status.idle":"2024-05-20T16:13:15.48002Z","shell.execute_reply.started":"2024-05-20T16:13:15.468842Z","shell.execute_reply":"2024-05-20T16:13:15.478929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FILE_PATH = ('/kaggle/input/rsna-breast-cancer-detection/'\n             'train_images/10006/1459541791.dcm')\n\nimg = read_xray(FILE_PATH, img_size=(512, 512))\n\nbatch = torch.tensor(img[None])\npatch_size = (16, 16)\nbatch_patches = patchify(batch, patch_size)\n\npatches = batch_patches[0]\nc, nh, nw, ph, pw = patches.shape\n\nplt.figure(figsize=(5, 5))\nplt.imshow(img[0], cmap=\"gray\")\nplt.axis(\"off\")\n\nplt.figure(figsize=(5, 5))\nfor i in range(nh):\n    for j in range(nw):\n        plt.subplot(nh, nw, i * nw + j + 1)\n        plt.imshow(patches[0, i, j], cmap=\"gray\")\n        plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:13:15.483Z","iopub.execute_input":"2024-05-20T16:13:15.483405Z","iopub.status.idle":"2024-05-20T16:14:02.851321Z","shell.execute_reply.started":"2024-05-20T16:13:15.48336Z","shell.execute_reply":"2024-05-20T16:14:02.85016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mlp(in_features, hidden_units, out_features):\n    \"\"\"\n    Returns a MLP head\n    \"\"\"\n    dims = [in_features] + hidden_units + [out_features]\n    layers = []\n    for dim1, dim2 in zip(dims[:-2], dims[1:-1]):\n        layers.append(nn.Linear(dim1, dim2))\n        layers.append(nn.ReLU())\n    layers.append(nn.Linear(dims[-2], dims[-1]))\n    return nn.Sequential(*layers)","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.85279Z","iopub.execute_input":"2024-05-20T16:14:02.853236Z","iopub.status.idle":"2024-05-20T16:14:02.860957Z","shell.execute_reply.started":"2024-05-20T16:14:02.853172Z","shell.execute_reply":"2024-05-20T16:14:02.859896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Img2Seq(nn.Module):\n    \"\"\"\n    This layers takes a batch of images as input and\n    returns a batch of sequences\n    \n    Shape:\n        input: (b, h, w, c)\n        output: (b, s, d)\n    \"\"\"\n    def __init__(self, img_size, patch_size, n_channels, d_model):\n        super().__init__()\n        self.patch_size = patch_size\n        self.img_size = img_size\n\n        nh, nw = img_size[0] // patch_size[0], img_size[1] // patch_size[1]\n        n_tokens = nh * nw\n\n        token_dim = patch_size[0] * patch_size[1] * n_channels\n        self.linear = nn.Linear(token_dim, d_model)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, d_model))\n        self.pos_emb = nn.Parameter(torch.randn(n_tokens, d_model))\n\n    def __call__(self, batch):\n        batch = patchify(batch, self.patch_size)\n\n        b, c, nh, nw, ph, pw = batch.shape\n\n        # Flattening the patches\n        batch = torch.permute(batch, [0, 2, 3, 4, 5, 1])\n        batch = torch.reshape(batch, [b, nh * nw, ph * pw * c])\n\n        batch = self.linear(batch)\n        cls = self.cls_token.expand([b, -1, -1])\n        emb = batch + self.pos_emb\n\n        return torch.cat([cls, emb], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.862339Z","iopub.execute_input":"2024-05-20T16:14:02.862743Z","iopub.status.idle":"2024-05-20T16:14:02.878952Z","shell.execute_reply.started":"2024-05-20T16:14:02.862707Z","shell.execute_reply":"2024-05-20T16:14:02.877864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViT(nn.Module):\n    def __init__(\n        self,\n        img_size,\n        patch_size,\n        n_channels,\n        d_model,\n        nhead,\n        dim_feedforward,\n        blocks,\n        mlp_head_units,\n        n_classes,\n    ):\n        super().__init__()\n        \"\"\"\n        Args:\n            img_size: Size of the image\n            patch_size: Size of the patch\n            n_channels: Number of image channels\n            d_model: The number of features in the transformer encoder\n            nhead: The number of heads in the multiheadattention models\n            dim_feedforward: The dimension of the feedforward network model in the encoder\n            blocks: The number of sub-encoder-layers in the encoder\n            mlp_head_units: The hidden units of mlp_head\n            n_classes: The number of output classes\n        \"\"\"\n        self.img2seq = Img2Seq(img_size, patch_size, n_channels, d_model)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model, nhead, dim_feedforward, activation=\"gelu\", batch_first=True\n        )\n        self.transformer_encoder = nn.TransformerEncoder(\n            encoder_layer, blocks\n        )\n        self.mlp = get_mlp(d_model, mlp_head_units, n_classes)\n        \n        self.output = nn.Sigmoid() if n_classes == 1 else nn.Softmax()\n\n    def __call__(self, batch):\n\n        batch = self.img2seq(batch)\n        batch = self.transformer_encoder(batch)\n        batch = batch[:, 0, :]\n        batch = self.mlp(batch)\n        output = self.output(batch)\n        return output","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.880532Z","iopub.execute_input":"2024-05-20T16:14:02.880905Z","iopub.status.idle":"2024-05-20T16:14:02.892099Z","shell.execute_reply.started":"2024-05-20T16:14:02.880866Z","shell.execute_reply":"2024-05-20T16:14:02.891142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = (512, 512)\npatch_size = (16, 16)\nn_channels = 1\nd_model = 1024\nnhead = 4\ndim_feedforward = 2048\nblocks = 8\nmlp_head_units = [1024, 512]\nn_classes = 1\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.893269Z","iopub.execute_input":"2024-05-20T16:14:02.893574Z","iopub.status.idle":"2024-05-20T16:14:02.948968Z","shell.execute_reply.started":"2024-05-20T16:14:02.89355Z","shell.execute_reply":"2024-05-20T16:14:02.947849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    \n    def __init__(self, df, img_path):\n        self.df = df\n        self.img_path = img_path\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        patient_id, image_id, cancer = self.df.iloc[idx][['patient_id', 'image_id', 'cancer']]\n        file = os.path.join(self.img_path, f'{patient_id}_{image_id}.png')\n        file = cv2.imread(file, cv2.COLOR_BGR2GRAY)\n        clahe = cv2.createCLAHE(clipLimit = 15, tileGridSize=[8, 8])\n        file = clahe.apply(file)\n        file = file / file.max()\n        X = torch.tensor(file[np.newaxis].astype('float32')).to(device)\n        y = torch.tensor([cancer]).float().to(device)\n        return X, y","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.950472Z","iopub.execute_input":"2024-05-20T16:14:02.950834Z","iopub.status.idle":"2024-05-20T16:14:02.960651Z","shell.execute_reply.started":"2024-05-20T16:14:02.950805Z","shell.execute_reply":"2024-05-20T16:14:02.95949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ncounts = df['cancer'].value_counts()\ndf['weights'] = df['cancer'].apply(lambda x: 1/counts[x])\n\ntrain_df, val_df = train_test_split(df, test_size=0.25, stratify=df['cancer'])","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:02.96483Z","iopub.execute_input":"2024-05-20T16:14:02.965173Z","iopub.status.idle":"2024-05-20T16:14:03.489761Z","shell.execute_reply.started":"2024-05-20T16:14:02.965146Z","shell.execute_reply":"2024-05-20T16:14:03.488649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '/kaggle/input/rsna-breast-cancer-512-pngs'\ntrain_samples = 1000\nval_samples = 500\n\ntrain_ds = RSNADataset(train_df, img_path)\nval_ds = RSNADataset(val_df, img_path)\n\ntrain_sampler = WeightedRandomSampler(train_df['weights'].values, train_samples)\ntrain_loader = DataLoader(train_ds, batch_size=8, sampler=train_sampler)\n\nval_sampler = WeightedRandomSampler(val_df['weights'].values, val_samples)\nval_loader = DataLoader(val_ds, batch_size=32, sampler=val_sampler)","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:03.490999Z","iopub.execute_input":"2024-05-20T16:14:03.491354Z","iopub.status.idle":"2024-05-20T16:14:03.498772Z","shell.execute_reply.started":"2024-05-20T16:14:03.491324Z","shell.execute_reply":"2024-05-20T16:14:03.497737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViT(\n    img_size = (512, 512),\n    patch_size = (16, 16),\n    n_channels = 1,\n    d_model = 1024,\n    nhead = 4,\n    dim_feedforward = 1024,\n    blocks = 8,\n    mlp_head_units = [512, 512],\n    n_classes = 1,\n).to(device)\n\noptimizer = Adam(model.parameters())\ncriterion = nn.BCELoss()\n\ntrainer = create_supervised_trainer(model, optimizer, criterion, device=device)\nval_metrics = {\n    \"bce\": Loss(criterion)\n}\nevaluator = create_supervised_evaluator(model, metrics=val_metrics, device=device)","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:03.500334Z","iopub.execute_input":"2024-05-20T16:14:03.500668Z","iopub.status.idle":"2024-05-20T16:14:03.93137Z","shell.execute_reply.started":"2024-05-20T16:14:03.50064Z","shell.execute_reply":"2024-05-20T16:14:03.930267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"log_interval = 10\nmax_epochs = 5\nbest_loss = float('inf')\n\nRunningAverage(output_transform=lambda x: x).attach(trainer, 'loss')\n\npbar = ProgressBar()\npbar.attach(trainer, ['loss'])\n    \n@trainer.on(Events.EPOCH_COMPLETED)\ndef log_validation_results(trainer):\n    global best_loss\n    evaluator.run(val_loader)\n    loss = evaluator.state.metrics['bce']\n    if loss < best_loss:\n        best_loss = loss\n        torch.save(model.state_dict(), 'best_model_vit.pt')\n    print(f\"Validation Results - Epoch: {trainer.state.epoch} Avg loss: {loss:.2f}\")\n    \noutput_state = trainer.run(train_loader, max_epochs=max_epochs)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:14:03.932677Z","iopub.execute_input":"2024-05-20T16:14:03.933085Z","iopub.status.idle":"2024-05-20T16:31:13.513028Z","shell.execute_reply.started":"2024-05-20T16:14:03.933049Z","shell.execute_reply":"2024-05-20T16:31:13.511684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom sklearn.metrics import f1_score, accuracy_score, precision_score\nimport numpy as np\n\n# Define the device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Assuming you have already defined and set up val_loader and your model architecture\n\n# Create an instance of your ViT model\nmodel = ViT(\n    img_size=(512, 512),\n    patch_size=(16, 16),\n    n_channels=1,\n    d_model=1024,\n    nhead=4,\n    dim_feedforward=1024,\n    blocks=8,\n    mlp_head_units=[512, 512],\n    n_classes=1\n).to(device)\n\n# Load the best model's weights\ncheckpoint_path = \"best_model_vit.pt\"\nmodel.load_state_dict(torch.load(checkpoint_path, map_location=device))\n\n# Set the model to evaluation mode\nmodel.eval()\n\n# Assuming val_loader is defined and properly set up\n# Evaluate the model on the validation data\nall_predictions = []\nall_labels = []\n\nwith torch.no_grad():\n    for data, labels in val_loader:\n        data = data.to(device)\n        labels = labels.to(device)\n        outputs = model(data)\n        predictions = torch.sigmoid(outputs).cpu().numpy()\n        all_predictions.extend(predictions)\n        all_labels.extend(labels.cpu().numpy())\n\n# Convert predictions to binary labels\nall_predictions = (np.array(all_predictions) > 0.5).astype(int)\nall_labels = np.array(all_labels)\n\n# Calculate F1 score\nf1 = f1_score(all_labels, all_predictions)\nprint(f\"F1 Score: {f1:.2f}\")\n\n# Calculate Accuracy\naccuracy = accuracy_score(all_labels, all_predictions)\nprint(f\"Accuracy: {accuracy:.2f}\")\n\n# Calculate Precision\nprecision = precision_score(all_labels, all_predictions)\nprint(f\"Precision: {precision:.2f}\")\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-20T16:31:13.51481Z","iopub.execute_input":"2024-05-20T16:31:13.515145Z","iopub.status.idle":"2024-05-20T16:31:41.377283Z","shell.execute_reply.started":"2024-05-20T16:31:13.515115Z","shell.execute_reply":"2024-05-20T16:31:41.37625Z"},"trusted":true},"execution_count":null,"outputs":[]}]}