{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4619805,"sourceType":"datasetVersion","datasetId":2688675}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\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-09-30T08:40:57.456552Z","iopub.execute_input":"2024-09-30T08:40:57.457009Z","iopub.status.idle":"2024-09-30T08:41:06.246046Z","shell.execute_reply.started":"2024-09-30T08:40:57.456967Z","shell.execute_reply":"2024-09-30T08:41:06.244545Z"},"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-09-30T08:41:06.248778Z","iopub.execute_input":"2024-09-30T08:41:06.249379Z","iopub.status.idle":"2024-09-30T08:41:06.258123Z","shell.execute_reply.started":"2024-09-30T08:41:06.249335Z","shell.execute_reply":"2024-09-30T08:41:06.256264Z"},"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-09-30T08:41:06.259919Z","iopub.execute_input":"2024-09-30T08:41:06.260462Z","iopub.status.idle":"2024-09-30T08:41:06.277399Z","shell.execute_reply.started":"2024-09-30T08:41:06.260383Z","shell.execute_reply":"2024-09-30T08:41:06.27607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FILE_PATH = ('/kaggle/input/rsna-breast-cancer-detection/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\")\n","metadata":{"execution":{"iopub.status.busy":"2024-09-30T08:41:06.279187Z","iopub.execute_input":"2024-09-30T08:41:06.279763Z","iopub.status.idle":"2024-09-30T08:41:53.725742Z","shell.execute_reply.started":"2024-09-30T08:41:06.279715Z","shell.execute_reply":"2024-09-30T08:41:53.724079Z"},"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-09-30T08:41:53.729382Z","iopub.execute_input":"2024-09-30T08:41:53.729835Z","iopub.status.idle":"2024-09-30T08:41:53.739517Z","shell.execute_reply.started":"2024-09-30T08:41:53.729793Z","shell.execute_reply":"2024-09-30T08:41:53.737717Z"},"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)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-30T08:41:53.741301Z","iopub.execute_input":"2024-09-30T08:41:53.741851Z","iopub.status.idle":"2024-09-30T08:41:53.756912Z","shell.execute_reply.started":"2024-09-30T08:41:53.741807Z","shell.execute_reply":"2024-09-30T08:41:53.755228Z"},"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-09-30T08:41:53.758557Z","iopub.execute_input":"2024-09-30T08:41:53.758938Z","iopub.status.idle":"2024-09-30T08:41:53.778385Z","shell.execute_reply.started":"2024-09-30T08:41:53.7589Z","shell.execute_reply":"2024-09-30T08:41:53.776848Z"},"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-09-30T08:41:53.78019Z","iopub.execute_input":"2024-09-30T08:41:53.780716Z","iopub.status.idle":"2024-09-30T08:41:53.797538Z","shell.execute_reply.started":"2024-09-30T08:41:53.780658Z","shell.execute_reply":"2024-09-30T08:41:53.796158Z"},"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-09-30T08:41:53.799357Z","iopub.execute_input":"2024-09-30T08:41:53.799901Z","iopub.status.idle":"2024-09-30T08:41:53.813328Z","shell.execute_reply.started":"2024-09-30T08:41:53.799848Z","shell.execute_reply":"2024-09-30T08:41:53.811372Z"},"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-09-30T08:41:53.815384Z","iopub.execute_input":"2024-09-30T08:41:53.815978Z","iopub.status.idle":"2024-09-30T08:41:54.400767Z","shell.execute_reply.started":"2024-09-30T08:41:53.815928Z","shell.execute_reply":"2024-09-30T08:41:54.399156Z"},"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-09-30T08:41:54.402525Z","iopub.execute_input":"2024-09-30T08:41:54.403056Z","iopub.status.idle":"2024-09-30T08:41:54.413488Z","shell.execute_reply.started":"2024-09-30T08:41:54.403003Z","shell.execute_reply":"2024-09-30T08:41:54.411714Z"},"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)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-30T08:41:54.41667Z","iopub.execute_input":"2024-09-30T08:41:54.417514Z","iopub.status.idle":"2024-09-30T08:41:54.728712Z","shell.execute_reply.started":"2024-09-30T08:41:54.417447Z","shell.execute_reply":"2024-09-30T08:41:54.727198Z"},"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-09-30T08:41:54.730582Z","iopub.execute_input":"2024-09-30T08:41:54.731035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resnet = models.resnet50(pretrained=True)\nin_features = resnet.fc.in_features\nresnet.fc = nn.Linear(in_features, 1)\n\nresnet_transforms= transforms.Compose([\n    transforms.Resize((228, 228)),\n    transforms.Lambda(lambda x: x.repeat([1, 3, 1, 1]))\n])\n\nclass MyResNet(nn.Module):\n    \n    def __init__(self, transforms, model):\n        super().__init__()\n        self.transforms = transforms\n        self.model = model\n        self.output = nn.Sigmoid()\n        \n    def forward(self, batch):\n        batch = self.transforms(batch)\n        batch = self.model(batch)\n        return self.output(batch)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}