{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Making use of the coordinates to come up with something useful","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch torchmetrics","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:15:34.275762Z","iopub.execute_input":"2024-07-09T09:15:34.276153Z","iopub.status.idle":"2024-07-09T09:15:54.764560Z","shell.execute_reply.started":"2024-07-09T09:15:34.276122Z","shell.execute_reply":"2024-07-09T09:15:54.762853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, gc, sys\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\ntqdm.pandas()\nimport glob\n\nimport math\nimport random\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport cv2\nimport PIL\nimport pydicom\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\n\nimport albumentations as A\n\nimport tensorflow as tf\n\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:25:40.571182Z","iopub.execute_input":"2024-07-09T09:25:40.571615Z","iopub.status.idle":"2024-07-09T09:25:40.579306Z","shell.execute_reply.started":"2024-07-09T09:25:40.571579Z","shell.execute_reply":"2024-07-09T09:25:40.578002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    print(\"seeding !!!\")\n    \ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n        \nseeding(42)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:12:42.147367Z","iopub.execute_input":"2024-07-09T08:12:42.148435Z","iopub.status.idle":"2024-07-09T08:12:42.158495Z","shell.execute_reply.started":"2024-07-09T08:12:42.148398Z","shell.execute_reply":"2024-07-09T08:12:42.157433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = Path(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\")\ntrain_main = pd.read_csv(DATA_PATH/\"train.csv\")\ntrain_labels = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\ntrain_desc = pd.read_csv(DATA_PATH/\"train_series_descriptions.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:02:37.483896Z","iopub.execute_input":"2024-07-09T08:02:37.484283Z","iopub.status.idle":"2024-07-09T08:02:37.587680Z","shell.execute_reply.started":"2024-07-09T08:02:37.484251Z","shell.execute_reply":"2024-07-09T08:02:37.586704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_desc = train_desc.merge(train_labels, how='inner', on=['study_id', 'series_id'])\ndf = label_desc.merge(train_main, how='inner', on='study_id')","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:08:45.327250Z","iopub.execute_input":"2024-07-09T08:08:45.327638Z","iopub.status.idle":"2024-07-09T08:08:45.382754Z","shell.execute_reply.started":"2024-07-09T08:08:45.327605Z","shell.execute_reply":"2024-07-09T08:08:45.381281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['image_path'] = (\n    f\"{str(DATA_PATH)}/train_images/\" + \n    df[\"study_id\"].astype(str) + \n    \"/\" + df[\"series_id\"].astype(str) + \n    \"/\" + df['instance_number'].astype(str) + \".dcm\"\n)\n\ncheck_path = lambda p: tf.io.gfile.exists(p)\ndf['exists'] = df['image_path'].progress_apply(check_path)\ndf['exists'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:25:33.956075Z","iopub.execute_input":"2024-07-09T08:25:33.957110Z","iopub.status.idle":"2024-07-09T08:25:39.076948Z","shell.execute_reply.started":"2024-07-09T08:25:33.957067Z","shell.execute_reply":"2024-07-09T08:25:39.075782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = data - np.min(data)\n        \n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:16:59.930251Z","iopub.execute_input":"2024-07-09T08:16:59.930767Z","iopub.status.idle":"2024-07-09T08:16:59.937440Z","shell.execute_reply.started":"2024-07-09T08:16:59.930732Z","shell.execute_reply":"2024-07-09T08:16:59.936207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"STUDY_ID = 4003253\nIMAGE_PATHS = glob.glob(f\"{DATA_PATH}/train_images/{str(STUDY_ID)}/*/*.dcm\")","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:17:00.990546Z","iopub.execute_input":"2024-07-09T08:17:00.992700Z","iopub.status.idle":"2024-07-09T08:17:01.001717Z","shell.execute_reply.started":"2024-07-09T08:17:00.992646Z","shell.execute_reply":"2024-07-09T08:17:01.000456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = load_dicom(IMAGE_PATHS[0])\nplt.imshow(img, cmap='gray');","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:17:02.308903Z","iopub.execute_input":"2024-07-09T08:17:02.309643Z","iopub.status.idle":"2024-07-09T08:17:02.624903Z","shell.execute_reply.started":"2024-07-09T08:17:02.309605Z","shell.execute_reply":"2024-07-09T08:17:02.623683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Rectangular masks","metadata":{}},{"cell_type":"code","source":"def create_rectangle_masks(img_path, coords, radius=64, show_plots=True):\n    center_coords = (int(coords['x']), int(coords['y']))\n    x0 = center_coords[0] - radius\n    x1 = center_coords[0] + radius\n    y0 = center_coords[1] - radius\n    y1 = center_coords[1] + radius\n    \n    # create a bounding box\n    top_left = (x0,y0)\n    bottom_right = (x1, y1)\n    \n    img = load_dicom(img_path)\n    mask = np.zeros_like(img).astype(np.uint8)\n    img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n    mask = cv2.rectangle(mask, top_left, bottom_right, color=255, thickness=cv2.FILLED)\n    if show_plots:\n        plt.figure(figsize=(7,7))\n        plt.subplot(131)\n        plt.imshow(img)\n        plt.axis(False)\n        plt.subplot(132)\n        plt.imshow(mask)\n        plt.axis(False)\n        plt.subplot(133)\n        plt.imshow(img)\n        plt.imshow(mask, alpha=0.5)\n        plt.axis(False)\n        plt.tight_layout()\n    else:\n        return img, mask\n\n\np = df.loc[10, 'image_path']\ncds = df.loc[10, ['x', 'y']]\ncreate_rectangle_masks(p, cds, radius=64)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:42:23.387243Z","iopub.execute_input":"2024-07-09T08:42:23.387685Z","iopub.status.idle":"2024-07-09T08:42:23.759908Z","shell.execute_reply.started":"2024-07-09T08:42:23.387651Z","shell.execute_reply":"2024-07-09T08:42:23.758779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_bounding_boxes(img_path, coords, radius=64, show_plots=True):\n    img, mask = create_rectangle_masks(img_path, coords, radius=radius, show_plots=False)\n    extracted_pixels = cv2.bitwise_and(img, img, mask=mask)\n    roi = np.where(extracted_pixels > 0)\n    x0 = roi[0].min()\n    x1 = roi[0].max()\n    y0 = roi[1].min()\n    y1 = roi[1].max()\n    \n    if show_plots:\n        plt.figure(figsize=(7,7))\n        plt.subplot(121)\n        plt.imshow(img)\n        plt.imshow(mask, alpha=0.5)\n        plt.axis(False)\n        plt.subplot(122)\n        plt.imshow(extracted_pixels[x0:x1, y0:y1])\n        plt.axis(False)\n        plt.tight_layout()\n        return None\n    else:\n        return extracted_pixels[x0:x1, y0:y1]\n    \np = df.loc[10, 'image_path']\ncds = df.loc[10, ['x', 'y']]\ncrop_bounding_boxes(img_path=p, coords=cds, radius=64, show_plots=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:42:48.954418Z","iopub.execute_input":"2024-07-09T08:42:48.955445Z","iopub.status.idle":"2024-07-09T08:42:49.346418Z","shell.execute_reply.started":"2024-07-09T08:42:48.955401Z","shell.execute_reply":"2024-07-09T08:42:49.344958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Circular masks","metadata":{}},{"cell_type":"code","source":"def create_circular_masks(img_path, coords, radius=64, show_plots=True):\n    center_coords = (int(coords['x']), int(coords['y']))\n    color = (255, 0, 0)\n    thickness = -1\n    img = load_dicom(img_path).squeeze()\n    mask = np.zeros_like(img).astype(np.uint8)\n    img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n    mask = cv2.normalize(mask, None, alpha=0, beta=255, norm_type=cv2.NORM_MINMAX, dtype=cv2.CV_8U)\n    mask = cv2.circle(mask.copy(), center_coords, radius, color, thickness)\n    if show_plots:\n        plt.figure(figsize=(7,7))\n        plt.subplot(131)\n        plt.imshow(img)\n        plt.axis(False)\n        plt.subplot(132)\n        plt.imshow(mask)\n        plt.axis(False)\n        plt.subplot(133)\n        plt.imshow(img)\n        plt.imshow(mask, alpha=0.5)\n        plt.axis(False)\n        plt.tight_layout()\n    else:\n        return img, mask\n\n\np = df.loc[10, 'image_path']\ncds = df.loc[10, ['x', 'y']]\ncreate_circular_masks(p, cds)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:44:44.752136Z","iopub.execute_input":"2024-07-09T08:44:44.752536Z","iopub.status.idle":"2024-07-09T08:44:45.124746Z","shell.execute_reply.started":"2024-07-09T08:44:44.752487Z","shell.execute_reply":"2024-07-09T08:44:45.123548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_circle_roi(img_path, coords, radius=64, show_plot=True):\n    img, mask = create_circular_masks(img_path=img_path, coords=coords, radius=radius, show_plots=False)\n    extracted_pixels = cv2.bitwise_and(img, img, mask=mask)\n    roi = np.where(extracted_pixels > 0)\n    x0 = roi[0].min()\n    x1 = roi[0].max()\n    y0 = roi[1].min()\n    y1 = roi[1].max()\n    \n    if show_plot:\n        plt.figure(figsize=(7,7))\n        plt.subplot(121)\n        plt.imshow(img)\n        plt.imshow(mask, alpha=0.5)\n        plt.axis(False)\n        plt.subplot(122)\n        plt.imshow(extracted_pixels[x0:x1, y0:y1])\n        plt.axis(False);\n        return None\n    else:\n        return extracted_pixels[x0:x1, y0:y1]\n    \np = df.loc[10, 'image_path']\ncds = df.loc[10, ['x', 'y']]\ncrop_circle_roi(p, cds)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T08:47:03.257709Z","iopub.execute_input":"2024-07-09T08:47:03.258110Z","iopub.status.idle":"2024-07-09T08:47:03.512668Z","shell.execute_reply.started":"2024-07-09T08:47:03.258070Z","shell.execute_reply":"2024-07-09T08:47:03.511484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Based on the visualizations, i think using rectangular mask bounding boxes will be very helpful.","metadata":{}},{"cell_type":"markdown","source":"# Create segmentation dataset","metadata":{}},{"cell_type":"code","source":"class SpineSegDataset(Dataset):\n    def __init__(self, data, transform=None):\n        self.data = data\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        img_path = self.data.loc[idx, 'image_path']\n        coords = self.data.loc[idx, ['x', 'y']]\n        img, mask = create_rectangle_masks(img_path, coords, radius=64, show_plots=False)\n        transformed = self.transform(image=img, mask=mask)\n        img, mask = transformed['image'], transformed['mask']\n        img = img.astype(np.float32).transpose(2,0,1) / 255.0\n        mask = mask.astype(np.float32) / 255.0\n        return {\"image\": img, \"mask\": mask}\n    \nts = A.Compose([\n    A.Resize(height=512, width=512),\n])\nds = SpineSegDataset(df, transform=ts)\ndls = DataLoader(ds, batch_size=8, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:28:57.989112Z","iopub.execute_input":"2024-07-09T09:28:57.989563Z","iopub.status.idle":"2024-07-09T09:28:57.999455Z","shell.execute_reply.started":"2024-07-09T09:28:57.989500Z","shell.execute_reply":"2024-07-09T09:28:57.998289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b = next(iter(dls))\nimages = b['image'].detach().cpu().numpy().transpose(0,2,3,1)\nmasks = b['mask'].detach().cpu().numpy()\n\nfig, axes = plt.subplots(1, 8, figsize=(12,12))\naxes = axes.flatten()\n\nfor i in range(8):\n    axes[i].imshow(images[i])\n    axes[i].imshow(masks[i], alpha=0.5)\n    axes[i].axis(False)\n    \nplt.tight_layout();","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:29:00.638229Z","iopub.execute_input":"2024-07-09T09:29:00.638647Z","iopub.status.idle":"2024-07-09T09:29:03.082195Z","shell.execute_reply.started":"2024-07-09T09:29:00.638612Z","shell.execute_reply":"2024-07-09T09:29:03.080901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model setup","metadata":{}},{"cell_type":"code","source":"CONFIG = dict(\n    in_channels = 3,\n    num_classes = 3,\n    lr = 1e-3,\n    img_size = 512,\n    encoder_weights = \"imagenet\",\n    encoder = \"resnext50_32x4d\",\n    device = torch.device(\"cuda:0\") if torch.cuda.is_available() else \"cpu\"\n)\n\n\ndef create_model(cfg):\n    model = smp.DeepLabV3Plus(\n        encoder_name=cfg['encoder'],\n        encoder_weights=cfg['encoder_weights'],\n        in_channels=cfg['in_channels'],\n        classes=cfg['num_classes'],\n        activation=None\n    )\n    return model.to(cfg['device'])\n\nmodel = create_model(CONFIG)\nmodel.eval()\noutputs = model(b['image'].to(torch.float32))","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:23:41.054696Z","iopub.execute_input":"2024-07-09T09:23:41.055117Z","iopub.status.idle":"2024-07-09T09:23:54.367804Z","shell.execute_reply.started":"2024-07-09T09:23:41.055077Z","shell.execute_reply":"2024-07-09T09:23:54.366653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\nactivated = F.softmax(input=outputs, dim=1)\npreds = torch.argmax(activated, dim=1)\nloss = criterion(outputs, b['mask'].to(torch.int64))\nloss\n# b['mask'].max()","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:29:28.509657Z","iopub.execute_input":"2024-07-09T09:29:28.510074Z","iopub.status.idle":"2024-07-09T09:29:28.962340Z","shell.execute_reply.started":"2024-07-09T09:29:28.510038Z","shell.execute_reply":"2024-07-09T09:29:28.961150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_mini_batches(batch, predicted_masks):\n    for image, gt_mask, pr_mask in zip(batch[\"image\"], batch[\"mask\"], predicted_masks):\n        plt.figure(figsize=(5, 5))\n\n        plt.subplot(1, 3, 1)\n        plt.imshow(image.permute(1,2,0).detach().cpu().numpy())  # convert CHW -> HWC\n        plt.title(\"Image\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 3, 2)\n        plt.imshow(gt_mask.detach().cpu().numpy().squeeze()) # just squeeze classes dim, because we have only one class\n        plt.title(\"Ground truth\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 3, 3)\n        plt.imshow(pr_mask.detach().cpu().numpy().squeeze()) # just squeeze classes dim, because we have only one class\n        plt.title(\"Prediction\")\n        plt.axis(\"off\")\n\n        plt.show()\n        \nvisualize_mini_batches(b, preds)","metadata":{"execution":{"iopub.status.busy":"2024-07-09T09:30:11.058188Z","iopub.execute_input":"2024-07-09T09:30:11.058652Z","iopub.status.idle":"2024-07-09T09:30:13.573597Z","shell.execute_reply.started":"2024-07-09T09:30:11.058618Z","shell.execute_reply":"2024-07-09T09:30:13.572282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}