{"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 os\nimport shutil\nfrom pathlib import Path\nimport nibabel as nib\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport torch\nimport torchvision\nimport imgaug\nimport imgaug.augmenters as iaa\nfrom imgaug.augmentables.bbs import BoundingBox\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom pytorch_lightning.loggers import TensorBoardLogger\nimport pydicom\nimport tarfile\nimport cv2\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-06T08:36:58.190110Z","iopub.execute_input":"2023-06-06T08:36:58.190570Z","iopub.status.idle":"2023-06-06T08:36:58.198778Z","shell.execute_reply.started":"2023-06-06T08:36:58.190538Z","shell.execute_reply":"2023-06-06T08:36:58.197681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing","metadata":{}},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsnacardiacdetectionlabels/rsna_heart_detection.csv\"\ndata = pd.read_csv(data_path)\ndata.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:36:58.201274Z","iopub.execute_input":"2023-06-06T08:36:58.201943Z","iopub.status.idle":"2023-06-06T08:36:58.230939Z","shell.execute_reply.started":"2023-06-06T08:36:58.201907Z","shell.execute_reply":"2023-06-06T08:36:58.229863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Copy Images","metadata":{}},{"cell_type":"code","source":"source_dir = \"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_images\"\ndestination_dir = \"./data\"\n\nsums, sums_squared = 0, 0\n\nfiles = [f\"{name}.dcm\" for name in data['name']]\n\nfor i, file in enumerate(files):\n    train_or_val = \"train\" if (i < 400) else \"val\"\n    dest_dir = os.path.join(destination_dir, train_or_val)\n\n    # Creating the destination directory if it doesn't exist\n    if not os.path.exists(dest_dir):\n        os.makedirs(dest_dir)\n\n    source_path = os.path.join(source_dir, file)\n    dest_path = os.path.join(dest_dir, file.split('.')[0])\n\n    dcm = pydicom.read_file(source_path)\n    dcm_arr = dcm.pixel_array\n    dcm_arr = (cv2.resize(dcm_arr, (224, 224))/255).astype(np.float32)\n\n    np.save(dest_path, dcm_arr)\n\n    normalizer = 224 * 224\n    if train_or_val == \"train\":\n        sums += np.sum(dcm_arr) / normalizer\n        sums_squared += np.sum(dcm_arr**2) / normalizer\n","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:36:58.232551Z","iopub.execute_input":"2023-06-06T08:36:58.232886Z","iopub.status.idle":"2023-06-06T08:37:01.877673Z","shell.execute_reply.started":"2023-06-06T08:36:58.232854Z","shell.execute_reply":"2023-06-06T08:37:01.876529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = Path('./data/train')\nval_path = Path('./data/val')\nsave_path = Path('./preprocessed')","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:01.880562Z","iopub.execute_input":"2023-06-06T08:37:01.881471Z","iopub.status.idle":"2023-06-06T08:37:01.888972Z","shell.execute_reply.started":"2023-06-06T08:37:01.881424Z","shell.execute_reply":"2023-06-06T08:37:01.885897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_rows = 2\nn_cols = 4\nfig, axes = plt.subplots(n_rows, n_cols, figsize=(8, 4))\nc = 0\nfor i in range(n_rows):\n    for j in range(n_cols):\n        pt_data = data.iloc[c]\n        pt_id = pt_data['name']\n        img_path = train_path/pt_id\n        img_path = img_path.with_suffix('.npy')\n        \n        img_arr = np.load(img_path)\n        \n        x = pt_data['x0']\n        y = pt_data['y0']\n        width = pt_data['w']\n        height = pt_data['h']\n        \n        axes[i][j].imshow(img_arr, cmap='bone')\n        rect = patches.Rectangle((x, y), width, height, linewidth=1, edgecolor='r', facecolor='none')\n        axes[i][j].add_patch(rect)\n        c += 1","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:01.894354Z","iopub.execute_input":"2023-06-06T08:37:01.895408Z","iopub.status.idle":"2023-06-06T08:37:03.622286Z","shell.execute_reply.started":"2023-06-06T08:37:01.895360Z","shell.execute_reply":"2023-06-06T08:37:03.621159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = sums / 400\nstd = np.sqrt((sums_squared / 400) - mean**2)\nmean, std","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:03.623544Z","iopub.execute_input":"2023-06-06T08:37:03.624968Z","iopub.status.idle":"2023-06-06T08:37:03.635028Z","shell.execute_reply.started":"2023-06-06T08:37:03.624929Z","shell.execute_reply":"2023-06-06T08:37:03.633980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class CardiacDataset(torch.utils.data.Dataset):\n    def __init__(self, path_to_labels_csv, root_path, augs):\n        self.labels = pd.read_csv(path_to_labels_csv)\n        self.root_path = Path(root_path)\n        self.patients = os.listdir(root_path)\n        self.augment = augs\n        \n    def __len__(self):\n        return len(self.patients)\n    \n    def __getitem__(self, idx):\n        patient = self.patients[idx]\n        data = self.labels[self.labels['name'] == patient.split('.')[0]]\n        \n        x_min = data['x0'].item()\n        y_min = data['y0'].item()\n        x_max = x_min + data['w'].item()\n        y_max = y_min + data['h'].item()\n        bbox = [x_min, y_min, x_max, y_max]\n        \n        file_path = self.root_path/patient\n        img = np.load(file_path).astype(np.float32)\n        \n        if self.augment:\n            bb = BoundingBox(x1=bbox[0], y1=bbox[1], x2=bbox[2], y2=bbox[3])\n            random_seed = torch.randint(0, int(1e+5), (1,)).item()\n            imgaug.seed(random_seed)\n            \n            img, aug_bbox = self.augment(image=img, bounding_boxes=bb)\n            bbox = aug_bbox[0][0], aug_bbox[0][1], aug_bbox[1][0], aug_bbox[1][1]\n        \n        img = (img - 0.494) / 0.252\n        img = torch.tensor(img).unsqueeze(0)\n        bbox = torch.tensor(bbox)\n        return img, bbox","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:03.636897Z","iopub.execute_input":"2023-06-06T08:37:03.639116Z","iopub.status.idle":"2023-06-06T08:37:03.668201Z","shell.execute_reply.started":"2023-06-06T08:37:03.639075Z","shell.execute_reply":"2023-06-06T08:37:03.666900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = iaa.Sequential([\n    iaa.GammaContrast(),\n    iaa.Affine(scale=(0.8, 1.2), rotate=(-10, 10), translate_px=(-10, 10))\n])","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:03.673576Z","iopub.execute_input":"2023-06-06T08:37:03.674774Z","iopub.status.idle":"2023-06-06T08:37:03.688886Z","shell.execute_reply.started":"2023-06-06T08:37:03.674733Z","shell.execute_reply":"2023-06-06T08:37:03.687470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CardiacDataset(data_path, train_path, train_transform)\nval_dataset = CardiacDataset(data_path, val_path, None)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:03.694065Z","iopub.execute_input":"2023-06-06T08:37:03.696952Z","iopub.status.idle":"2023-06-06T08:37:03.716117Z","shell.execute_reply.started":"2023-06-06T08:37:03.696898Z","shell.execute_reply":"2023-06-06T08:37:03.715184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing Dataset","metadata":{}},{"cell_type":"code","source":"img, bbox = val_dataset[0]\n\nfig, ax = plt.subplots(1, 1, figsize=(8, 4))\n\nax.imshow(img[0], cmap='bone')\nrect = patches.Rectangle((bbox[0], bbox[1]), bbox[2]-bbox[0], bbox[3]-bbox[1], linewidth=1, edgecolor='r', facecolor='none')\nax.add_patch(rect)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:03.721021Z","iopub.execute_input":"2023-06-06T08:37:03.723473Z","iopub.status.idle":"2023-06-06T08:37:04.223220Z","shell.execute_reply.started":"2023-06-06T08:37:03.723432Z","shell.execute_reply":"2023-06-06T08:37:04.222129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 16\nnum_workers = 2\n\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size, num_workers=num_workers, shuffle=True)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size, num_workers=num_workers, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.224671Z","iopub.execute_input":"2023-06-06T08:37:04.225201Z","iopub.status.idle":"2023-06-06T08:37:04.234017Z","shell.execute_reply.started":"2023-06-06T08:37:04.225106Z","shell.execute_reply":"2023-06-06T08:37:04.232864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class CardiacDetectionModel(pl.LightningModule):\n    def __init__(self):\n        super().__init__()\n        \n        self.model = torchvision.models.resnet18(weights=False)\n        self.model.conv1 = torch.nn.Conv2d(\n            in_channels=1, out_channels=64, kernel_size=7,\n            stride=2, padding=3, bias=False)\n        self.model.fc = torch.nn.Linear(in_features=512, out_features=4)\n        \n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=1e-4)\n        self.loss_fn = torch.nn.MSELoss()\n        \n    def forward(self, data):\n        return self.model(data)\n    \n    def training_step(self, batch, batch_idx):\n        img, label = batch\n        label = label.float()\n        pred = self(img)\n        loss = self.loss_fn(pred, label)\n        \n        self.log(\"Train Loss\", loss)\n        if batch_idx % 50 == 0:\n            self.log_images(img.cpu(), pred.cpu(), label.cpu(), \"Train\")\n            \n        return loss\n        \n        \n    def validation_step(self, batch, batch_idx):\n        img, label = batch\n        label = label.float()\n        pred = self(img)\n        loss = self.loss_fn(pred, label)\n        \n        self.log(\"Val Loss\", loss)\n        if batch_idx % 10 == 0:\n            self.log_images(img.cpu(), pred.cpu(), label.cpu(), \"Val\")\n            \n        return loss\n    \n    def log_images(self, x_ray, pred, label, name):\n        results = []\n        \n        for i in range(4):\n            coord_labels = label[i]\n            coord_pred = pred[i]\n            \n            img = ((x_ray[i] * 0.252) + 0.494).numpy()[0]\n            \n            x0, y0 = coord_labels[0].int().item(), coord_labels[1].int().item()\n            x1, y1 = coord_labels[2].int().item(), coord_labels[3].int().item()\n            img = cv2.rectangle(img, (x0, y0), (x1, y1), (0, 0, 0), 2)\n            \n            \n            x0, y0 = coord_pred[0].int().item(), coord_pred[1].int().item()\n            x1, y1 = coord_pred[2].int().item(), coord_pred[3].int().item()\n            img = cv2.rectangle(img, (x0, y0), (x1, y1), (1, 1, 1), 2)\n            \n            results.append(torch.tensor(img).unsqueeze(0))\n        \n        grid = torchvision.utils.make_grid(results, 2)\n        self.logger.experiment.add_image(name, grid, self.global_step)\n        \n    def configure_optimizers(self):\n        return [self.optimizer]","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.235859Z","iopub.execute_input":"2023-06-06T08:37:04.236280Z","iopub.status.idle":"2023-06-06T08:37:04.257457Z","shell.execute_reply.started":"2023-06-06T08:37:04.236243Z","shell.execute_reply":"2023-06-06T08:37:04.256331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CardiacDetectionModel()\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.261431Z","iopub.execute_input":"2023-06-06T08:37:04.261706Z","iopub.status.idle":"2023-06-06T08:37:04.439766Z","shell.execute_reply.started":"2023-06-06T08:37:04.261681Z","shell.execute_reply":"2023-06-06T08:37:04.438484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"checkpoint_callback = ModelCheckpoint(monitor=\"Val Loss\", save_top_k=10, mode=\"min\")","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.444636Z","iopub.execute_input":"2023-06-06T08:37:04.444991Z","iopub.status.idle":"2023-06-06T08:37:04.451932Z","shell.execute_reply.started":"2023-06-06T08:37:04.444961Z","shell.execute_reply":"2023-06-06T08:37:04.450765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = pl.Trainer(logger=TensorBoardLogger('./logs'), log_every_n_steps=1, callbacks=checkpoint_callback, max_epochs=100)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.453592Z","iopub.execute_input":"2023-06-06T08:37:04.454034Z","iopub.status.idle":"2023-06-06T08:37:04.516881Z","shell.execute_reply.started":"2023-06-06T08:37:04.453996Z","shell.execute_reply":"2023-06-06T08:37:04.515914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:04.519344Z","iopub.execute_input":"2023-06-06T08:37:04.520313Z","iopub.status.idle":"2023-06-06T08:37:15.874538Z","shell.execute_reply.started":"2023-06-06T08:37:04.520270Z","shell.execute_reply":"2023-06-06T08:37:15.873479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation","metadata":{}},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# model = model.load_from_checkpoint('./logs/lightning_logs/version_0/checkpoints/epoch=99-step=2500.ckpt')\n# model.eval()\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:42.447956Z","iopub.execute_input":"2023-06-06T08:37:42.448368Z","iopub.status.idle":"2023-06-06T08:37:42.475297Z","shell.execute_reply.started":"2023-06-06T08:37:42.448340Z","shell.execute_reply":"2023-06-06T08:37:42.474230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nlabels = []\n\nwith torch.no_grad():\n    for data, label in val_dataset:\n        data = data.to(device).float().unsqueeze(0)\n        # removing the batch channel dim\n        pred = model(data)[0].cpu()\n        preds.append(pred)\n        labels.append(label)\n        \npreds = torch.stack(preds)\nlabels = torch.stack(labels)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:43.523111Z","iopub.execute_input":"2023-06-06T08:37:43.524102Z","iopub.status.idle":"2023-06-06T08:37:44.046556Z","shell.execute_reply.started":"2023-06-06T08:37:43.524066Z","shell.execute_reply":"2023-06-06T08:37:44.045499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"abs(preds - labels).mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:46.195522Z","iopub.execute_input":"2023-06-06T08:37:46.195985Z","iopub.status.idle":"2023-06-06T08:37:46.205223Z","shell.execute_reply.started":"2023-06-06T08:37:46.195950Z","shell.execute_reply":"2023-06-06T08:37:46.204185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualize predictions","metadata":{}},{"cell_type":"code","source":"n_rows = 2\nn_cols = 4\nfig, axes = plt.subplots(n_rows, n_cols, figsize=(12, 6))\nidx = 0\nfor i in range(n_rows):\n    for j in range(n_cols):\n        img, label = val_dataset[idx]\n        pred = preds[idx]\n        \n        axes[i][j].imshow(img[0], cmap='bone')\n        heart = patches.Rectangle((pred[0], pred[1]), pred[2]-pred[0], pred[3]-pred[1], edgecolor='r', facecolor='none')\n        axes[i][j].add_patch(heart)\n        idx += 1","metadata":{"execution":{"iopub.status.busy":"2023-06-06T08:37:16.497602Z","iopub.status.idle":"2023-06-06T08:37:16.498073Z","shell.execute_reply.started":"2023-06-06T08:37:16.497843Z","shell.execute_reply":"2023-06-06T08:37:16.497864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}