{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## SenNetworks: Hacking the Human Vasculature in 3D\n\n### Exploratory notebook using Segformer and SAM\n\n#### Notebook Author: Andrew Kettle\n\n#### Note that this notebook should NOT be submitted, it is for exploration only due to licensing.","metadata":{}},{"cell_type":"code","source":"# Packages\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nimport torch\nimport torchvision.transforms.functional as TF\nimport torch.nn.functional as F\nimport pandas as pd\nimport os\nimport shutil\nimport transformers # Hugging Face for Segformer models\nimport torch.optim.lr_scheduler as lr_scheduler\n\nfrom tqdm import tqdm\nfrom torchinfo import summary\nfrom sklearn.model_selection import train_test_split\nfrom segment_anything import SamAutomaticMaskGenerator, SamPredictor, sam_model_registry","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset\nThe following few cells implement the data ingestion functions for the kidney dataset. The dataset and dataloader are both implemented using the base PyTorch classes. ","metadata":{}},{"cell_type":"code","source":"# Dataset that is passed to dataloader\nclass SenDataset(torch.utils.data.Dataset):\n\n\tdef __init__(self, run_type, feature_extractor):\n\t\tself.run_type = run_type\n\t\tself.feature_extractor = feature_extractor\n\t\tself.image_paths = []\n\t\tself.label_paths = []\n\n\tdef setPaths(self, image_paths, label_paths):\n\t\tself.image_paths = image_paths\n\n\t\t# Guard\n\t\tif self.run_type != 'test':\n\t\t\tself.label_paths = label_paths\n\n\tdef setImagePath(self, image_paths):\n\t\tself.image_paths = image_paths\n\t\n\tdef importTif(self, path, gray: bool):\n\t\timage = cv2.imread(path) # Read\n\t\tif gray:\n\t\t\t# Grab a single channel for label\n\t\t\timage = cv2.imread(path)[:,:,0] # R=G=B, so index can be [0,1,2]\n\n\t\treturn image\n\n\tdef __len__(self):\n\t\treturn len(self.image_paths)\n\n\tdef __getitem__(self, idx):\n\n\t\tif self.run_type != 'test':\n\t\t\t# Grab image and label from disk\n\t\t\timage = self.importTif(self.image_paths[idx], False)\n\t\t\tlabel = self.importTif(self.label_paths[idx], True)\n\n\t\t\t# Scale label, do this because Segformer uses 255 as an \"ignore index\"\n\t\t\tlabel = (label / 255).astype('uint8') # Scale to [0,1]\n\n\t\t\t# Prepare inputs for segformer using SegformerImageProcessor\n\t\t\tencoded_inputs = self.feature_extractor(image, label, return_tensors='pt')\n\t\t\tfor k,v in encoded_inputs.items():\n\t\t\t\tencoded_inputs[k].squeeze_() # remove batch axis\n\n\t\t\t# Only need image and labels for our implementation\n\t\t\timage = encoded_inputs['pixel_values']\n\t\t\tlabel = encoded_inputs['labels']\n\n\t\t\treturn image, label\n\n\t\telse:\n\t\t\t# Grab image from disk (no label provided for test set)\n\t\t\timage = self.importTif(self.image_paths[idx], False)\n\n\t\t\t# Prepare inputs for segformer using SegformerImageProcessor\n\t\t\tencoded_inputs = self.feature_extractor(image, return_tensors='pt')\n\t\t\tfor k,v in encoded_inputs.items():\n\t\t\t\tencoded_inputs[k].squeeze() # remove batch axis\n\n\t\t\timage = encoded_inputs['pixel_values']\n\t\t\treturn image, ''\n\n# Rename items in the dataset so there is no overlap (only run this once!)\ndef renameDataset():\n\troot_dir = os.path.join(os.getcwd(), \"data\")\n\ttrain_dir = os.path.join(root_dir, \"train\")\n\n\tpaths = os.listdir(train_dir)\n\tpaths = [os.path.join(train_dir, p) for p in paths]\n\n\t# Iterate over each subset\n\tfor p in paths:\t\n\t\t\n\t\timage_path = os.path.join(p, \"images\")\n\t\tlabel_path = os.path.join(p, \"labels\")\n\t\timage_names = os.listdir(image_path)\n\t\tlabel_names = os.listdir(label_path)\n\n\t\t# Name to append for each file (splitting is platform specific)\n\t\tappend_name = p.split('/')[-1]\n\n\t\trevised_image_names = [append_name + '_' + name for name in image_names]\n\t\trevised_label_names = [append_name + '_' + name for name in label_names]\n\n\t\t# Rename images\n\t\tfor iname, upiname in zip(image_names, revised_image_names):\n\t\t\tshutil.move(os.path.join(image_path, iname), os.path.join(image_path, upiname))\n\t\t# Rename labels\n\t\tfor lname, uplname in zip(label_names, revised_label_names):\n\t\t\tshutil.move(os.path.join(label_path, lname), os.path.join(label_path, uplname))\n\n# Populate numpy array with necessary dataset information\ndef buildDataset(root_dir: str, run_type: str):\n\t\"\"\" Populate the image and label arrays\n\n\tArgs:\n\t\troot_dir (str): Path to root dir (where training or testing may take place)\n\t\trun_type (str): either 'train' or 'test'\n\t\"\"\"\n\n\timage_paths = []\n\tlabel_paths = []\n\n\t# Logic for training or testing\n\tif run_type == 'train':\n\t\tworking_dir = os.path.join(root_dir, \"train\")\n\telif run_type == 'test':\n\t\tworking_dir = os.path.join(root_dir, \"test\")\n\telse:\n\t\tprint(\"Unsupported run_type in the buildDataset function\")\n\t\tassert(False) # Quick hack to exit execution\n\n\tpaths = os.listdir(working_dir)\n\tpaths = [os.path.join(working_dir, p) for p in paths]\n\n\tfor p in paths:\n\n\t\timage_path = os.path.join(p, \"images\")\n\t\timage_names = [os.path.join(image_path, name) for name in os.listdir(image_path)]\n\t\timage_paths.append(np.array(image_names))\n\n\t\tif run_type == 'train':\n\t\t\tlabel_path = os.path.join(p, \"labels\")\n\t\t\tlabel_names = [os.path.join(label_path, name) for name in os.listdir(label_path)]\n\t\t\tlabel_paths.append(np.array(label_names))\n\t\n\timage_paths = np.concatenate(image_paths)\n\n\tif run_type == 'train':\n\t\tlabel_paths = np.concatenate(label_paths)\n\n\treturn image_paths, label_paths","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Cell for calculating the image mean and standard deviation of the dataset.\n# We resize to 512x512 so that everything is standardized\n\ndef calcMeanAndStd():\n\n\t# Change paths if necessary!!\n\troot_dir = os.path.join(os.getcwd(), \"data\")\n\ttrain_dir = os.path.join(root_dir, \"train\")\n\n\tpaths = os.listdir(train_dir)\n\tpaths = [os.path.join(train_dir, p) for p in paths]\n\n\tus = 0\n\tstds = 0\n\tits = 0\n\n\t# Iterate over each subset\n\tfor p in paths:\t\n\t\t\n\t\timage_path = os.path.join(p, \"images\")\n\t\timage_names = os.listdir(image_path)\n\n\t\tfor name in image_names:\n\t\t\timage = cv2.imread(os.path.join(image_path, name))\n\t\t\timage = cv2.resize(image, (512, 512))\n\t\t\tus = us + np.mean(image) \n\t\t\tstds = stds + np.std(image) \n\t\t\tits+=1\n\n\tprint(us / float(its))\n\tprint(stds / float(its))\n\n#calcMeanAndStd()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ingest data, store in numpy arrays\nXtr, Ytr = buildDataset(os.path.join(os.getcwd(), 'data'), 'train')\nprint(f\"Xtr.shape: {Xtr.shape} Ytr.shape: {Ytr.shape}\")\n\nXte, _ = buildDataset(os.path.join(os.getcwd(), 'data'), 'test')\nprint(f\"Xte.shape: {Xte.shape}\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Cleaning\n\nObserve that the training set has fewer labels than images (Depends on which kidney folders are being used). Since we are data rich, we will elect to toss out the samples that don't have labels for the sake of simplicity.","metadata":{}},{"cell_type":"code","source":"# Trim file extensions\na = np.array([x.split('.')[0].split('/')[-1] for x in Xtr])\nb = np.array([y.split('.')[0].split('/')[-1] for y in Ytr])\n\n# Find overlap\nc = np.where(np.intersect1d(a, b))\n\n# Drop extra items from training set\nXtr = Xtr[c]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Split data into training and testing using 80/20 split","metadata":{}},{"cell_type":"code","source":"# Split dataset into training and validation\nXtr, Xval, Ytr, Yval = train_test_split(Xtr, Ytr, test_size=0.2, shuffle=True)\nprint(f\"Xtr: {Xtr.shape}, Ytr: {Ytr.shape}\")\nprint(f\"Xval: {Xval.shape}, Yval: {Yval.shape}\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Create dataset and dataloader","metadata":{}},{"cell_type":"code","source":"# Numbers extracted from mean and std calculation function\n# Using default imagenet mean and stddev because those are working better\nimage_mean = 103.84\nimage_std = 7.92\n\n# Create preprocessing routine for segformer (sometimes called feature_extractor)\nfeature_extractor_train = transformers.SegformerImageProcessor()\nfeature_extractor_test = transformers.SegformerImageProcessor(do_normalize=False) # Test images shouldn't know about training\n\n# Create datasets and set paths\ntrds = SenDataset('train', feature_extractor_train)\ntrds.setPaths(Xtr, Ytr)\n\nvalds = SenDataset('validation', feature_extractor_train)\nvalds.setPaths(Xval, Yval)\n\nteds = SenDataset('test', feature_extractor_test)\nteds.setImagePath(Xte)\n\n# Maximum size that an RTX3080 with 12GB VRAM could handle.\nbatch_size = 8\n\ntrain_loader = torch.utils.data.DataLoader(trds, batch_size=batch_size, num_workers=4)\nvalid_loader = torch.utils.data.DataLoader(valds, batch_size=batch_size, num_workers=4)\ntest_loader = torch.utils.data.DataLoader(teds, batch_size=1) # Batch size 1 for testing","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Cell for sanity checking dataset methods\n\n# Grab item from training dataset\nimage, label = trds.__getitem__(1)\n\n# Ensure shapes are correct\nprint(f\"Image Size: {image.shape}\")\nprint(f\"Label Size: {label.shape}\")\n\n# Convert to numpy for display\nimage = image.numpy()\nlabel = label.numpy()\n\n# Check that labels are in correct range\nprint(f\"Unique Label Values: {np.unique(label)}\") # Should be 0 or 1\n\n# Reshape for display\nimage = np.moveaxis(image, 0, -1)\nimage = (image * 255).astype('uint8') # Image after norm\n\n# Display image\nfig, ax = plt.subplots(1, 2, figsize=(15, 15))\nax[0].imshow(image)\nax[1].imshow(label, cmap='gray')\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Implement Segformer on the Vascular Dataset","metadata":{}},{"cell_type":"markdown","source":"### Loss Function Implementations for training\n\nThese loss functions are not used in the current iteration of the notebook. I was experimenting with them earlier, but I have modified other parts of the notebook so these should be verified before use.","metadata":{}},{"cell_type":"code","source":"# Pytorch Dice Loss Implmentation Modified from: https://www.kaggle.com/code/bigironsphere/loss-function-library-keras-pytorch\n# DiceLoss for cross entropy loss (assumes num_classes > 1)\nclass DiceLoss(torch.nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        # Compute logits from sigmoid\n        logits = torch.sigmoid(inputs)       \n\n        # Compute softmax from logits\n        sig = torch.sigmoid(logits)\n        softmax = F.softmax(sig, dim=1) # Compute softmax to get probs\n\n        probs, _ = torch.max(softmax, dim=1)\n\n        # Seg cts in [0,1], masks discrete one-hot with labels [0,1]\n        intersection = (probs * targets).sum()\n        dice = (2. * intersection + 1) / (probs.sum() + targets.sum() + 1)\n        \n        return 1 - dice\n\n# Weighted cross entropy loss using unnormalized logits, not being used\n# TODO: Bug check this function\nclass WeightedCE(torch.nn.Module):\n    def __init__(self, weights):\n        super(WeightedCE, self).__init__()\n        self.weights = weights\n        self.criterion = torch.nn.BCEWithLogitsLoss(pos_weight=torch.tensor(weights[1]))\n\n    def forward(self, logits, targets):\n\n        # Copmpute sigmoid and softmax explicitly\n        sig = torch.sigmoid(logits)\n        softmax = F.softmax(logits, dim=1) # Compute softmax to get probs\n        probs, _ = torch.max(softmax, dim=1)\n\n        probs = probs.flatten()\n        targets = torch.flatten(targets)\n\n        loss = (-1 * self.weights[1] * targets * torch.log(probs)) - (self.weights[0] * (1 - targets) * torch.log(1 - targets))\n        return torch.mean(loss)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Adopted from here: https://github.com/huggingface/transformers/issues/15819 (Thanks Niels!)\n\nclass SegFormerCustom(transformers.SegformerPreTrainedModel):\n\tdef __init__(self, config, pretrained):\n\t\tsuper().__init__(config)\n\t\tself.segformer = transformers.SegformerForSemanticSegmentation.from_pretrained(pretrained, num_labels=config.num_labels)\n\t\tself.pretrained_path = pretrained\n\t\tself.num_labels = config.num_labels\n\t\t#self.loss_func = WeightedCE(torch.tensor([1, 1]))\n\n\tdef forward(self, pixel_values, labels=None):\n\t\t# get raw logits from segformer model\n\t\toutputs = self.segformer(pixel_values, labels)\n\t\treturn outputs.logits, outputs.loss\n\n\t\t# If you would like to use a custom loss function, uncomment the code below and remove the return above\n\t\t# Works by bypassing the segformer loss and caculating loss using the \n\t\t# unnormalized logits directly\n\t\t\"\"\"\n\t\tlogits = outputs.logits\n\t\tloss = None\n\t\tif labels is not None: # Training\n\t\t\tupsampled_logits = F.interpolate(logits,\n\t\t\t\t\tsize=pixel_values.size()[2:], # (height, width)\n\t\t\t\t\tmode='bilinear',\n\t\t\t\t\talign_corners=False)\n\n\t\t\t# Compute DiceBCE loss using logits from segformer\n\t\t\tloss = self.loss_func(upsampled_logits, labels.float())\n\n\t\treturn logits, loss\n\t\t\"\"\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create configuration for Segformer Model\nconfig = transformers.SegformerConfig(num_labels=1) # 1 label = binary classification\n\n# Grab pretrained model\n# Use mit-b0 right now because it is the most lightweight\nmodel = SegFormerCustom(config, 'nvidia/mit-b0')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Move model to GPU\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu' # Check for device\nmodel = model.to(device) # Send model to device before creating optimizer\n\n# Optimizer\noptimizer = torch.optim.AdamW(model.parameters(), lr=0.000006) # Segformer Default LR: 0.000006\n\n# Learning Rate Scheduler. Not in use right now because of transfer learning\nscheduler = lr_scheduler.MultiStepLR(optimizer, milestones=[20, 40], gamma=0.1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training Loop\ntotal_steps = len(train_loader)\nepochs = 60\n\ntrain_loss_avg = []\nval_loss_avg = []\n\n# Train network\nfor ep in range(epochs):\n\n\tper_epoch_loss_train = []\n\tper_epoch_loss_val = []\n\n\t# Training\n\tfor idx, (images, masks) in enumerate(tqdm(train_loader)):\n\t\t# Convert vars to GPU\n\t\timages = images.float().to(device)\n\t\tmasks = masks.type(torch.LongTensor).to(device)\n\n\t\toutput = model(pixel_values=images, labels=masks)\n\t\tloss = output[1]\n\t\tloss.backward()\n\n\t\toptimizer.step()\n\t\toptimizer.zero_grad()\n\n\t\t# Save batches loss\n\t\tper_epoch_loss_train.append(loss.item())\n\n\tprint(f\"Epoch: {ep}: Training Loss: {np.mean(np.array(per_epoch_loss_train))}\")\n\n\t# Validation\n\tfor idx, (images, masks) in enumerate(tqdm(valid_loader)):\n\t\t# Convert vars to GPU\n\t\timages = images.float().to(device)\n\t\tmasks = masks.type(torch.LongTensor).to(device)\n\n\t\toutput = model(pixel_values=images, labels=masks)\n\t\tloss = output[1]\n\n\t\t# Save batches loss\n\t\tper_epoch_loss_val.append(loss.item())\n\n\tscheduler.step()\n\t\n\ttrain_loss_avg.append(np.mean(np.array(per_epoch_loss_train)))\n\tval_loss_avg.append(np.mean(np.array(per_epoch_loss_val)))\n\n# Empty GPU memory\nwith torch.no_grad():\n\ttorch.cuda.empty_cache()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Empty GPU memory\nwith torch.no_grad():\n\ttorch.cuda.empty_cache()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot results\nplt.plot(range(0, len(train_loss_avg)), train_loss_avg, label=\"Training Loss\", color=\"blue\")\nplt.plot(range(0, len(val_loss_avg)), val_loss_avg, label=\"Validation Loss\", color=\"orange\")\nplt.title(\"Loss curves during Segformer Training\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.show();","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save entire trained model to disk\ntorch.save(model, 'segformer-b0-kidney-1.pt')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load a segformer model for testing (if desired)\nmodel = torch.load('model/segformer-b0-kidney-1.pt')\n\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu' # Check for device\nmodel = model.to(device) # Send model to device before creating optimizer","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test Trained Segformer Model on validation data","metadata":{}},{"cell_type":"code","source":"# Decode outputs from Segformer's unnormalized logits\n\ndef calcDice(probs, mask):\n\t# Seg cts in [0,1], masks discrete one-hot with labels [0,1]\n\tintersection = (probs * mask).sum()\n\tdice = (2. * intersection + 1) / (probs.sum() + mask.sum() + 1)\n\n\treturn dice\n\ndef decodeSegformer(logits, mask=None):\n\t\"\"\" Decode the segformer model\n\n\tArgs:\n\t\tlogits (_type_): raw logits from segformer calculation\n\t\"\"\"\n\n\tupsampled_logits = F.interpolate(logits,\n\t\tsize=image.size()[2:], # (height, width)\n\t\tmode='bilinear',\n\t\talign_corners=False)\n\n\tprobs = torch.sigmoid(upsampled_logits)\n\n\tif mask is not None:\n\t\tdice = calcDice(probs, mask)\n\telse:\n\t\tdice = -1\n\n\t# pred above .5 = 1, below = 0\n\tseg = np.round(probs)\n\n\treturn seg, dice\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Run-Length Encode and Decode From Carvana on Kaggle\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport time\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n \ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test the trained segformer model\n\nwith torch.no_grad():\n\n    model.eval()\n\n    segs = []\n    dices = []\n    images = []\n    masks = []\n\n    # Get random indices for evaluating\n    rng = np.random.default_rng()\n    rands = rng.integers(low=0, high=valds.__len__(), size=3)        \n    \n    for i in rands:\n        image, mask = valds.__getitem__(i)\n        images.append(image)\n        masks.append(mask)\n\n    for image, mask in zip(images, masks):\n        image = image.to(device)\n        mask = mask.to(device)\n\n        image = torch.unsqueeze(image, 0)\n        mask = torch.unsqueeze(mask, 0)\n        outputs = model(pixel_values=image, labels=mask)\n        logits = outputs[0].cpu()\n        mask = mask.cpu()\n\n        seg, dice = decodeSegformer(logits, mask)\n\n        dices.append(dice)\n        segs.append(seg)\n\n    # Convert images and labels back to device\n    images = [x.cpu().numpy() for x in images]\n    masks =  [x.cpu().numpy() for x in masks]\n\n    # Show examples\n    fig, ax = plt.subplots(3, 3, figsize = (10, 10))\n    fig.tight_layout() # Make plot look better\n\n    ax[0,0].set_title('Images')\n    ax[0,1].set_title('Predictions')\n    ax[0,2].set_title('Labels')\n\n    # Images\n    ax[0,0].imshow(images[0][0])\n    ax[1,0].imshow(images[1][0])\n    ax[2,0].imshow(images[2][0])\n\n    # Predictions\n    ax[0,1].imshow(np.squeeze(segs[0]))\n    ax[0,1].text(5, 25, f'dice: {dices[0].item():.3f}', bbox={'facecolor': 'white', 'pad': 2})\n    ax[1,1].imshow(np.squeeze(segs[1]))\n    ax[1,1].text(5, 25, f'dice: {dices[1].item():.3f}', bbox={'facecolor': 'white', 'pad': 2})\n    ax[2,1].imshow(np.squeeze(segs[2]))\n    ax[2,1].text(5, 25, f'dice: {dices[2].item():.3f}', bbox={'facecolor': 'white', 'pad': 2})\n\n    # Labels\n    ax[0,2].imshow(masks[0])\n    ax[1,2].imshow(masks[1])\n    ax[2,2].imshow(masks[2])\n\n    plt.show()\n\n    \"\"\"\n    for r in segs:\n        seg = torch.squeeze(r).numpy()\n        mask_rle = rle_encode(seg)\n        print(\"Mask_rle: \", mask_rle)\n        dec = rle_decode(mask_rle, seg.shape)\n    \"\"\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Submission Code","metadata":{}},{"cell_type":"code","source":"# Test the trained segformer model\n\nwith torch.no_grad():\n\n    model.eval()\n\n    segs = []\n    dices = []\n    images = []\n    masks = []\n\n    for i in range(0, 6):\n        image, _ = teds.__getitem__(i)\n        images.append(image)\n\n    for image in images:\n        image = image.to(device)\n\n        #image = torch.unsqueeze(image, 0)\n        outputs = model(pixel_values=image)\n        logits = outputs[0].cpu()\n\n        seg, _ = decodeSegformer(logits)\n\n        segs.append(seg)\n\n    # Convert images and labels back to device\n    images = [x.cpu().numpy() for x in images]\n\n    # Show examples\n    fig, ax = plt.subplots(3, 2, figsize = (10, 10))\n    fig.tight_layout() # Make plot look better\n\n    ax[0,0].set_title('Images')\n    ax[0,1].set_title('Predictions')\n\n    images = np.squeeze(images, 1)\n\n    # Images\n    ax[0,0].imshow(images[0][0])\n    ax[1,0].imshow(images[1][0])\n    ax[2,0].imshow(images[2][0])\n\n    # Predictions\n    ax[0,1].imshow(np.squeeze(segs[0]))\n    ax[1,1].imshow(np.squeeze(segs[1]))\n    ax[2,1].imshow(np.squeeze(segs[2]))\n\n    plt.show()\n\n    for r in segs:\n        seg = torch.squeeze(r).numpy()\n        mask_rle = rle_encode(seg)\n        if mask_rle == \"\":\n            mask_rle = \"1 0\"\n        print(\"Mask_rle: \", mask_rle)\n        dec = rle_decode(mask_rle, seg.shape)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extra: Testing Segment-Anything (SAM) Zero-Shot Segmentation Capabilities","metadata":{}},{"cell_type":"code","source":"# Import SAM model\n\ncheckpoint = os.path.join(os.path.join(os.getcwd(), 'model'), 'sam_vit_h_4b8939.pth')\n\ndevice = 'cuda:0' if torch.cuda.is_available else 'cpu'\n\nsam = sam_model_registry['vit_h'](checkpoint=checkpoint)\nsam.to(device=device) # Run on GPU\nmask_generator = SamAutomaticMaskGenerator(sam)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Grab a few items from the training set, using data_trans for submission\nk1_path = os.path.join(os.path.join(os.path.join(os.path.join(os.getcwd(), 'data'), 'train'), 'kidney_1_dense'), 'images')\nnames = os.listdir(k1_path)\nimages = []\n\nimages.append(cv2.imread(os.path.join(k1_path, names[0])))\nimages.append(cv2.imread(os.path.join(k1_path, names[1])))\nimages.append(cv2.imread(os.path.join(k1_path, names[2])))\n\n# Make prediction\ntotalmasks = []\nfor image in images:\n\tmasks = mask_generator.generate(image)\n\ttotalmasks.append(np.array(masks))\n\n# Empty GPU memory\nwith torch.no_grad():\n\ttorch.cuda.empty_cache()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# From segment-anything automatic_mask_generator_example.ipynb in the FAIR Github repository\ndef show_anns(anns):\n    if len(anns) == 0:\n        return\n    sorted_anns = sorted(anns, key=(lambda x: x['area']), reverse=True)\n\n    img = np.ones((sorted_anns[0]['segmentation'].shape[0], sorted_anns[0]['segmentation'].shape[1], 4))\n\n    img[:,:,3] = 0\n    for ann in sorted_anns:\n        m = ann['segmentation']\n        color_mask = np.concatenate([np.random.random(3), [0.35]])\n        img[m] = color_mask\n    return img\n\n# Show examples\nfig, ax = plt.subplots(2, 3, figsize = (10, 8))\nfig.tight_layout() # Make plot look better\n\nax[0,0].set_title('Images')\nax[1,0].set_title('Predictions')\n\n# Images\nax[0,0].imshow(images[0])\nax[0,1].imshow(images[1])\nax[0,2].imshow(images[2])\n\n# Predictions\nax[1,0].imshow(show_anns(totalmasks[0]))\nax[1,0].set_autoscale_on(False)\nax[1,1].imshow(show_anns(totalmasks[1]))\nax[1,1].set_autoscale_on(False)\nax[1,2].imshow(show_anns(totalmasks[2]))\nax[1,2].set_autoscale_on(False)\n\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Zero-shot SAM Performance\n\nStraight out of the box, SAM performs poorly for our vascular dataset challenge. Looking at the middle column, it is evident that some segmentation is taking place, but it isn't relevant to the vascular problem. The list of issues is summarized below. \n\n1. SAM needs to be prompted (bounding box, text, etc) which isn't included in our dataset. We could potentially augment bounding boxes, but the dataset doesn't seem that well suited to rectangular boxes. \n2. The structure of biological images differs heavily from \"natural images\". Examples: Different collection modaltiies (X-Ray, ultrasound, etc), No color features (images are grayscale upscaled to RGB)\n3. SAM is performing instance segmentation, but we have a binary segmentation problem\n4. SAM does not have knowledge of the vascular problem we are trying to solve.\n\nConclusion:  \nProblem one is the biggest issue for the given dataset. Generating bounding boxes seems difficult given the structure of the segmentation maps. Problems 2-4 could be solved by fine tuning SAM. ","metadata":{}}]}