{"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":"markdown","source":"# Ratio\nsee [here](https://www.kaggle.com/competitions/hubmap-hacking-the-human-vasculature/discussion/419143)\n1. Train a yolov8 model to identify all the annotated types\n2. Fine tune a SAM Model\n3. predict on an image with YOLO & SAM","metadata":{}},{"cell_type":"markdown","source":"# Train Val Test","metadata":{}},{"cell_type":"markdown","source":"## Generate BBoxes for all labels","metadata":{}},{"cell_type":"code","source":"IN_TRAIN_IMGS_PATH = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train\"\nTRAIN_IMGS_PATH = ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil","metadata":{"execution":{"iopub.status.busy":"2023-07-24T16:58:20.269081Z","iopub.execute_input":"2023-07-24T16:58:20.270165Z","iopub.status.idle":"2023-07-24T16:58:20.274649Z","shell.execute_reply.started":"2023-07-24T16:58:20.270129Z","shell.execute_reply":"2023-07-24T16:58:20.273573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/labels","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:31:31.095052Z","iopub.execute_input":"2023-07-24T17:31:31.095462Z","iopub.status.idle":"2023-07-24T17:31:32.155452Z","shell.execute_reply.started":"2023-07-24T17:31:31.09543Z","shell.execute_reply":"2023-07-24T17:31:32.153851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = {\"blood_vessel\" : 0, \"glomerulus\" : 1, \"unsure\" : 2}","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:36:32.748824Z","iopub.execute_input":"2023-07-24T17:36:32.749223Z","iopub.status.idle":"2023-07-24T17:36:32.756645Z","shell.execute_reply.started":"2023-07-24T17:36:32.749191Z","shell.execute_reply":"2023-07-24T17:36:32.755658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def xyxy_to_cxcywh(bbox, img):\n    '''\n    convert from xyxy notation \n    to \n    normalized center_x center_y, width, height notation\n    origin is top left\n    c.f. [here](https://user-images.githubusercontent.com/26833433/91506361-c7965000-e886-11ea-8291-c72b98c25eec.jpg)\n    '''\n    xmin, ymin, xmax, ymax = bbox[0], bbox[1], bbox[2], bbox[3]\n    \n    assert xmax >= xmin\n    assert ymax >= ymin\n    \n    dw = 1.0 / img.shape[0]\n    dh = 1.0 / img.shape[1]\n    \n    W = float(xmax - xmin)\n    H = float(ymax - ymin)\n    \n    cx = (xmin + xmax) / 2.0\n    cy = (ymin + ymax) / 2.0\n    \n    # normalize\n    cx *= dw\n    cy *= dh\n    H *= dh\n    W *= dw\n    \n    return [cx,cy,W,H]","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:36:33.438463Z","iopub.execute_input":"2023-07-24T17:36:33.43883Z","iopub.status.idle":"2023-07-24T17:36:33.450742Z","shell.execute_reply.started":"2023-07-24T17:36:33.438798Z","shell.execute_reply":"2023-07-24T17:36:33.449652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nGenerate Yolo labels\nFor each image:\n    1. get bounding boxes\n    2. convert bboxes to YOLO format\n    3. save each in txt file with corresponding class (fmt: class cx cy W H)\n'''\ndf = pd.read_json(\"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\", lines=True)\n\nfor i in tqdm(range(len(df))):\n    img_id = df.iloc[i]['id']\n    out_str = \"\"\n    for j in range(len(df.iloc[i]['annotations'])): # each label for this image\n        label = df.iloc[i]['annotations'][j]\n        \n        # get bounding box\n        ground_truth_mask = np.array(label['coordinates'])\n        mask = cv2.fillPoly(np.zeros((IMG_SIZE, IMG_SIZE)), [ground_truth_mask.reshape((-1,1,2))], color=(255,0,0))\n        bbox = get_bounding_box(mask)\n        cxcy_bbox = xyxy_to_cxcywh(bbox, mask)\n        \n        line = str(classes[label['type']]) + \" \"\n        line += ' '.join(map(str, cxcy_bbox))\n        \n        out_str += line + \"\\n\"\n        \n    # write to file\n    with open(f\"/kaggle/working/labels/{img_id}.txt\", \"w\") as f:\n        f.write(out_str)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:37:07.186539Z","iopub.execute_input":"2023-07-24T17:37:07.18765Z","iopub.status.idle":"2023-07-24T17:37:31.407928Z","shell.execute_reply.started":"2023-07-24T17:37:07.187607Z","shell.execute_reply":"2023-07-24T17:37:31.406574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Split Val and Training Datasets","metadata":{}},{"cell_type":"code","source":"!rm -rf /kaggle/working/train/labels/train\n!rm -rf /kaggle/working/train/labels/val\n!rm -rf /kaggle/working/train/images/train\n!rm -rf /kaggle/working/train/images/val","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:37:31.410534Z","iopub.execute_input":"2023-07-24T17:37:31.411255Z","iopub.status.idle":"2023-07-24T17:37:35.615269Z","shell.execute_reply.started":"2023-07-24T17:37:31.411215Z","shell.execute_reply":"2023-07-24T17:37:35.613811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /kaggle/working/train/labels/train\n!mkdir -p /kaggle/working/train/labels/val\n!mkdir -p /kaggle/working/train/images/train\n!mkdir -p /kaggle/working/train/images/val","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:37:35.617433Z","iopub.execute_input":"2023-07-24T17:37:35.61812Z","iopub.status.idle":"2023-07-24T17:37:39.710584Z","shell.execute_reply.started":"2023-07-24T17:37:35.618079Z","shell.execute_reply":"2023-07-24T17:37:39.709158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = os.listdir(\"/kaggle/working/labels\")\n# imgs = os.listdir(\"/kaggle/input/hubmap-hacking-the-human-vasculature/train/\")\ntrain = labels[:200]\nval = labels[200:300]\n\nfor im in tqdm(train):\n    im = im.split(\".\")[0]\n    shutil.copyfile(f\"/kaggle/input/hubmap-hacking-the-human-vasculature/train/{im}.tif\", f\"/kaggle/working/train/images/train/{im}.tif\")\n    shutil.move(f\"/kaggle/working/labels/{im}.txt\", f\"/kaggle/working/train/labels/train/{im}.txt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:37:39.713823Z","iopub.execute_input":"2023-07-24T17:37:39.714559Z","iopub.status.idle":"2023-07-24T17:37:40.123733Z","shell.execute_reply.started":"2023-07-24T17:37:39.71452Z","shell.execute_reply":"2023-07-24T17:37:40.122743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val = labels[200:300]\nfor im in tqdm(val):\n    im = im.split(\".\")[0]\n    shutil.copyfile(f\"/kaggle/input/hubmap-hacking-the-human-vasculature/train/{im}.tif\", f\"/kaggle/working/train/images/val/{im}.tif\")\n    shutil.move(f\"/kaggle/working/labels/{im}.txt\", f\"/kaggle/working/train/labels/val/{im}.txt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:38:17.061842Z","iopub.execute_input":"2023-07-24T17:38:17.062226Z","iopub.status.idle":"2023-07-24T17:38:17.295472Z","shell.execute_reply.started":"2023-07-24T17:38:17.062196Z","shell.execute_reply":"2023-07-24T17:38:17.294529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Segmentation","metadata":{}},{"cell_type":"code","source":"!pip install -q git+https://github.com/huggingface/transformers.git &2>/dev/null","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:30.301018Z","iopub.execute_input":"2023-07-24T17:53:30.301865Z","iopub.status.idle":"2023-07-24T17:53:31.349949Z","shell.execute_reply.started":"2023-07-24T17:53:30.301829Z","shell.execute_reply":"2023-07-24T17:53:31.348494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q monai","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:31.356134Z","iopub.execute_input":"2023-07-24T17:53:31.35863Z","iopub.status.idle":"2023-07-24T17:53:47.364385Z","shell.execute_reply.started":"2023-07-24T17:53:31.358591Z","shell.execute_reply":"2023-07-24T17:53:47.363055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision \nimport torchvision.transforms as transforms\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nfrom PIL import Image, ImageDraw\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:47.367199Z","iopub.execute_input":"2023-07-24T17:53:47.367928Z","iopub.status.idle":"2023-07-24T17:53:50.815159Z","shell.execute_reply.started":"2023-07-24T17:53:47.367888Z","shell.execute_reply":"2023-07-24T17:53:50.814197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:50.818541Z","iopub.execute_input":"2023-07-24T17:53:50.819465Z","iopub.status.idle":"2023-07-24T17:53:50.850031Z","shell.execute_reply.started":"2023-07-24T17:53:50.819427Z","shell.execute_reply":"2023-07-24T17:53:50.849137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SIZE = 512\nBATCH_SIZE = 2\nLEARNING_RATE = 1e-5","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:50.853992Z","iopub.execute_input":"2023-07-24T17:53:50.854295Z","iopub.status.idle":"2023-07-24T17:53:50.86509Z","shell.execute_reply.started":"2023-07-24T17:53:50.854247Z","shell.execute_reply":"2023-07-24T17:53:50.86419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_bounding_box(ground_truth_map, pturb=20):\n  # get bounding box from mask\n  y_indices, x_indices = np.where(ground_truth_map > 0)\n  x_min, x_max = np.min(x_indices), np.max(x_indices)\n  y_min, y_max = np.min(y_indices), np.max(y_indices)\n  # add perturbation to bounding box coordinates\n  H, W = ground_truth_map.shape\n  x_min = max(0, x_min - np.random.randint(0, pturb))\n  x_max = min(W, x_max + np.random.randint(0, pturb))\n  y_min = max(0, y_min - np.random.randint(0, pturb))\n  y_max = min(H, y_max + np.random.randint(0, pturb))\n  bbox = [x_min, y_min, x_max, y_max]\n\n  return bbox","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:50.866306Z","iopub.execute_input":"2023-07-24T17:53:50.866877Z","iopub.status.idle":"2023-07-24T17:53:50.876512Z","shell.execute_reply.started":"2023-07-24T17:53:50.866839Z","shell.execute_reply":"2023-07-24T17:53:50.875701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubMapDataset(Dataset):\n    def __init__(self, img_dir, polygons_file, processor, num_samples=None):\n        self.root_dir = img_dir\n        self.images = os.listdir(img_dir)\n        \n        self.labels = self._separate_labels(polygons_file)\n        self.processor = processor\n        self.num_samples\n    \n    def _separate_labels(self, polygons_file):\n        '''Assign each label to its own row in a DF'''\n        df = pd.read_json(polygons_file, lines=True)\n        new_df_ls = []\n        for i in range(len(df)):\n            curr_id = df.iloc[i]['id']\n            for j in range(len(df.iloc[i]['annotations'])):\n                new_entry = df.iloc[i]['annotations'][j]\n                new_entry['id'] = curr_id\n                new_df_ls.append(new_entry)\n\n        self.labels = pd.DataFrame(new_df_ls)\n        # convert coordinates to np array\n        self.labels.coordinates = [np.asarray(coors).squeeze() for coors in self.labels.coordinates]\n        \n        ########################################################################################\n        if self.num_samples is None:\n            return self.labels\n        else:\n            assert type(self.num_samples) is int\n            return self.labels.sample(self.num_samples, random_state=12345) \n        ########################################################################################\n        \n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, idx):\n        label = self.labels.iloc[idx]\n        curr_id = label['id']\n        img = cv2.imread(os.path.join(self.root_dir, f\"{curr_id}.tif\"))\n        \n        # get bounding box prompt\n        ground_truth_mask = label['coordinates']\n        mask = cv2.fillPoly(np.zeros((img.shape[0], img.shape[1])), [ground_truth_mask.reshape((-1,1,2))], color=(255,0,0))\n        # mask_r = cv2.resize(mask, (N_IMG_SIZE, N_IMG_SIZE))\n            \n        prompt = get_bounding_box(mask)\n        \n        # print(f\"Mask: {mask.shape}\")\n        #print(f\"OG Img: {img.shape}\")\n        #print(f\"Bounding Box: {prompt}\")\n\n        # prepare image and prompt for the model\n        inputs = self.processor(img, input_boxes=[[prompt]], return_tensors=\"pt\")\n\n        # remove batch dimension which the processor adds by default\n        inputs = {k:v.squeeze(0) for k,v in inputs.items()}\n\n        # add ground truth segmentation\n        inputs[\"ground_truth_mask\"] = mask #cv2.resize(mask, (N_IMG_SIZE, N_IMG_SIZE))\n        \n        #print(\"#####\\n# inputs: #\\n######\")\n        ##print(f\"Mask: {inputs['ground_truth_mask'].shape}\")\n        #print(f\"New img: {inputs['pixel_values'].shape}\")\n        #print(f\"BBox: {inputs['input_boxes']}\")\n\n        return inputs\n     ","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:50.877854Z","iopub.execute_input":"2023-07-24T17:53:50.878666Z","iopub.status.idle":"2023-07-24T17:53:50.893194Z","shell.execute_reply.started":"2023-07-24T17:53:50.878632Z","shell.execute_reply":"2023-07-24T17:53:50.892235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import SamProcessor\nfrom torch.utils.data import DataLoader\n\nprocessor = SamProcessor.from_pretrained(\"facebook/sam-vit-base\")","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:50.894604Z","iopub.execute_input":"2023-07-24T17:53:50.894951Z","iopub.status.idle":"2023-07-24T17:53:59.644083Z","shell.execute_reply.started":"2023-07-24T17:53:50.894919Z","shell.execute_reply":"2023-07-24T17:53:59.643125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dset = HubMapDataset(\"/kaggle/input/hubmap-hacking-the-human-vasculature/train\", \"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\", processor)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:53:59.645454Z","iopub.execute_input":"2023-07-24T17:53:59.646901Z","iopub.status.idle":"2023-07-24T17:54:07.556148Z","shell.execute_reply.started":"2023-07-24T17:53:59.646864Z","shell.execute_reply":"2023-07-24T17:54:07.551504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dset, batch_size=BATCH_SIZE, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:07.56279Z","iopub.execute_input":"2023-07-24T17:54:07.563135Z","iopub.status.idle":"2023-07-24T17:54:07.582299Z","shell.execute_reply.started":"2023-07-24T17:54:07.563105Z","shell.execute_reply":"2023-07-24T17:54:07.576285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(train_dataloader))\nfor k,v in batch.items():\n  print(k,v.shape)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:07.584583Z","iopub.execute_input":"2023-07-24T17:54:07.584894Z","iopub.status.idle":"2023-07-24T17:54:07.904023Z","shell.execute_reply.started":"2023-07-24T17:54:07.584865Z","shell.execute_reply":"2023-07-24T17:54:07.903151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import SamModel \n\nmodel = SamModel.from_pretrained(\"facebook/sam-vit-base\")\n\n# make sure we only compute gradients for mask decoder\nfor name, param in model.named_parameters():\n  if name.startswith(\"vision_encoder\") or name.startswith(\"prompt_encoder\"):\n    param.requires_grad_(False)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:07.908192Z","iopub.execute_input":"2023-07-24T17:54:07.910398Z","iopub.status.idle":"2023-07-24T17:54:12.68441Z","shell.execute_reply.started":"2023-07-24T17:54:07.910348Z","shell.execute_reply":"2023-07-24T17:54:12.68341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train the SAM Model","metadata":{}},{"cell_type":"code","source":"from torch.optim import Adam\nimport monai\n\nfrom statistics import mean\nimport torch\nfrom torch.nn.functional import threshold, normalize","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:12.685765Z","iopub.execute_input":"2023-07-24T17:54:12.686205Z","iopub.status.idle":"2023-07-24T17:54:15.761892Z","shell.execute_reply.started":"2023-07-24T17:54:12.686171Z","shell.execute_reply":"2023-07-24T17:54:15.760912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Note: Hyperparameter tuning could improve performance here\noptimizer = Adam(model.mask_decoder.parameters(), lr=LEARNING_RATE, weight_decay=0)\n\nseg_loss = monai.losses.DiceCELoss(sigmoid=True, squared_pred=True, reduction='mean')","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:15.763295Z","iopub.execute_input":"2023-07-24T17:54:15.76374Z","iopub.status.idle":"2023-07-24T17:54:15.771292Z","shell.execute_reply.started":"2023-07-24T17:54:15.763701Z","shell.execute_reply":"2023-07-24T17:54:15.770159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 20 # 100 is recommended\n\nmodel.to(device)\nmodel.train()\nfor epoch in range(num_epochs):\n    epoch_losses = []\n    for batch in tqdm(train_dataloader):\n        # forward pass\n        outputs = model(pixel_values=batch[\"pixel_values\"].to(device),\n                      input_boxes=batch[\"input_boxes\"].to(device),\n                      multimask_output=False)\n\n        # compute loss\n        predicted_masks = outputs.pred_masks.squeeze(1)\n        ground_truth_masks = batch[\"ground_truth_mask\"].squeeze().float().to(device)\n\n        ground_truth_masks = nn.functional.interpolate(ground_truth_masks.unsqueeze(0),\n                size=(256, 256),\n                mode='bilinear',\n                align_corners=False)\n\n        assert predicted_masks.squeeze().shape == ground_truth_masks.squeeze().shape, f\"{predicted_masks.shape} does not match {ground_truth_masks.shape}\"\n\n        loss = seg_loss(predicted_masks.squeeze(), ground_truth_masks.squeeze())\n        # print(loss)\n\n        # backward pass (compute gradients of parameters w.r.t. loss)\n        optimizer.zero_grad()\n        loss.backward()\n\n        # optimize\n        optimizer.step()\n        epoch_losses.append(loss.item())\n\n    print(f'EPOCH: {epoch}')\n    print(f'Mean loss: {mean(epoch_losses)}')\n","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:54:53.793603Z","iopub.execute_input":"2023-07-24T17:54:53.793968Z","iopub.status.idle":"2023-07-24T18:18:26.741131Z","shell.execute_reply.started":"2023-07-24T17:54:53.793938Z","shell.execute_reply":"2023-07-24T18:18:26.740167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save!\ntorch.save(model.state_dict(), \"/kaggle/working/out.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:37:52.064018Z","iopub.execute_input":"2023-07-24T18:37:52.064451Z","iopub.status.idle":"2023-07-24T18:37:52.590272Z","shell.execute_reply.started":"2023-07-24T18:37:52.064419Z","shell.execute_reply":"2023-07-24T18:37:52.589211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(\"out.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:39:23.0411Z","iopub.execute_input":"2023-07-24T18:39:23.041784Z","iopub.status.idle":"2023-07-24T18:39:23.048307Z","shell.execute_reply.started":"2023-07-24T18:39:23.04175Z","shell.execute_reply":"2023-07-24T18:39:23.047327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"from pylab import rcParams\n\nrcParams['figure.figsize'] = 18, 8","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:44:57.118139Z","iopub.execute_input":"2023-07-24T18:44:57.118618Z","iopub.status.idle":"2023-07-24T18:44:57.130604Z","shell.execute_reply.started":"2023-07-24T18:44:57.118581Z","shell.execute_reply":"2023-07-24T18:44:57.129645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BBOX_STROKE=7\nidx = 68\nimg_id = train_dset.labels.iloc[idx]['id']\nprint(img_id)\nprint(train_dset.labels.iloc[idx]['type'])\n\nimg_dir = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train\"\nimg = cv2.imread(f\"{img_dir}/{img_id}.tif\")\nimg_proc = train_dset[idx]['pixel_values'].permute(1,2,0) # permute Tensor to HxWxC\n\nbxs = train_dset[idx]['input_boxes'].squeeze().numpy().astype(np.int32)\nimg_bb = cv2.rectangle(cv2.resize(img, (1024,1024)), (bxs[0], bxs[1]), (bxs[2], bxs[3]), (0,255,0), BBOX_STROKE)\nimg_proc_bb = cv2.rectangle(cv2.resize(img_proc.numpy(), (1024,1024)), (bxs[0], bxs[1]), (bxs[2], bxs[3]), (0,255,0), BBOX_STROKE)\n\ngt_mask = cv2.fillPoly(np.zeros((IMG_SIZE,IMG_SIZE)), [train_dset.labels.iloc[idx].coordinates.reshape((-1,1,2))], color=(255,0,0))","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:44:58.801429Z","iopub.execute_input":"2023-07-24T18:44:58.8018Z","iopub.status.idle":"2023-07-24T18:44:58.990059Z","shell.execute_reply.started":"2023-07-24T18:44:58.801771Z","shell.execute_reply":"2023-07-24T18:44:58.989091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get box prompt based on ground truth segmentation map\nprompt = get_bounding_box(gt_mask)\n\n# prepare image + box prompt for the model\ninputs = processor(img, input_boxes=[[prompt]], return_tensors=\"pt\").to(device)\nfor k,v in inputs.items():\n  print(k,v.shape)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:00.896861Z","iopub.execute_input":"2023-07-24T18:45:00.897611Z","iopub.status.idle":"2023-07-24T18:45:00.970881Z","shell.execute_reply.started":"2023-07-24T18:45:00.89757Z","shell.execute_reply":"2023-07-24T18:45:00.969902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show Image and processed image (normalized, resized, etc)\nf, axarr = plt.subplots(1,2)\naxarr[0].imshow(img)\naxarr[1].imshow(img_proc)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:01.778295Z","iopub.execute_input":"2023-07-24T18:45:01.779014Z","iopub.status.idle":"2023-07-24T18:45:02.984154Z","shell.execute_reply.started":"2023-07-24T18:45:01.778972Z","shell.execute_reply":"2023-07-24T18:45:02.982469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show bounding box on \nf, axarr = plt.subplots(1,2)\naxarr[0].imshow(img_bb)\naxarr[1].imshow(img_proc_bb)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:03.450735Z","iopub.execute_input":"2023-07-24T18:45:03.451104Z","iopub.status.idle":"2023-07-24T18:45:04.799205Z","shell.execute_reply.started":"2023-07-24T18:45:03.451074Z","shell.execute_reply":"2023-07-24T18:45:04.798323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\n\nwith torch.no_grad():\n    outputs = model(**inputs, multimask_output=False)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:04.800764Z","iopub.execute_input":"2023-07-24T18:45:04.801669Z","iopub.status.idle":"2023-07-24T18:45:05.061573Z","shell.execute_reply.started":"2023-07-24T18:45:04.801633Z","shell.execute_reply":"2023-07-24T18:45:05.060578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resize_pred=False\npredicted_masks = outputs.pred_masks.squeeze(1)\nground_truth_masks = torch.from_numpy(gt_mask).squeeze().float().to(device)\nprint(f\"Predicted Mask shape: {predicted_masks.shape}\")\nprint(f\"GT mask shape: {ground_truth_masks.unsqueeze(0).unsqueeze(0).shape}\")\n\nif resize_pred:\n    predicted_masks = nn.functional.interpolate(predicted_masks,\n            size=(IMG_SIZE, IMG_SIZE),\n            mode='bilinear',\n            align_corners=False)\nelse:\n    ground_truth_masks = nn.functional.interpolate(ground_truth_masks.unsqueeze(0).unsqueeze(0),\n        size=(256,256),\n        mode='bilinear',\n        align_corners=False)\n\nassert predicted_masks.squeeze().shape == ground_truth_masks.squeeze().shape, f\"{predicted_masks.shape} does not match {ground_truth_masks.shape}\"\n\nloss = seg_loss(predicted_masks.squeeze(), ground_truth_masks.squeeze())\nloss","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:09.059397Z","iopub.execute_input":"2023-07-24T18:45:09.059767Z","iopub.status.idle":"2023-07-24T18:45:09.08279Z","shell.execute_reply.started":"2023-07-24T18:45:09.059739Z","shell.execute_reply":"2023-07-24T18:45:09.081433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# apply sigmoid\nmedsam_seg_prob = torch.sigmoid(outputs.pred_masks.squeeze(1))\n# convert soft mask to hard mask\nmedsam_seg_prob = medsam_seg_prob.cpu().numpy().squeeze()\nmedsam_seg = (medsam_seg_prob > 0.5).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:13.099492Z","iopub.execute_input":"2023-07-24T18:45:13.099854Z","iopub.status.idle":"2023-07-24T18:45:13.105669Z","shell.execute_reply.started":"2023-07-24T18:45:13.099824Z","shell.execute_reply":"2023-07-24T18:45:13.104594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# GT mask vs. predicted mask\nf, axarr = plt.subplots(1,2)\naxarr[0].imshow(img)\naxarr[0].imshow(gt_mask, alpha=0.4)\naxarr[1].imshow(cv2.resize(img_proc.numpy(), (256,256)))\naxarr[1].imshow(medsam_seg, alpha=0.4)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T18:45:15.346964Z","iopub.execute_input":"2023-07-24T18:45:15.347858Z","iopub.status.idle":"2023-07-24T18:45:16.616233Z","shell.execute_reply.started":"2023-07-24T18:45:15.347815Z","shell.execute_reply":"2023-07-24T18:45:16.615397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# YoloV8","metadata":{}},{"cell_type":"markdown","source":"## Train Yolo Model","metadata":{}},{"cell_type":"code","source":"!pip install -q ultralytics","metadata":{"execution":{"iopub.status.busy":"2023-07-24T19:00:26.340671Z","iopub.execute_input":"2023-07-24T19:00:26.341381Z","iopub.status.idle":"2023-07-24T19:00:38.289686Z","shell.execute_reply.started":"2023-07-24T19:00:26.34133Z","shell.execute_reply":"2023-07-24T19:00:38.28827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ultralytics import YOLO","metadata":{"execution":{"iopub.status.busy":"2023-07-24T19:00:38.293533Z","iopub.execute_input":"2023-07-24T19:00:38.293855Z","iopub.status.idle":"2023-07-24T19:00:40.01224Z","shell.execute_reply.started":"2023-07-24T19:00:38.293826Z","shell.execute_reply":"2023-07-24T19:00:40.01109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BLOOD_VESSEL = 0\nGLOMERULUS = 1\nUNSURE = 2\n\nIMG_SIZE = 512","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:04:43.106401Z","iopub.execute_input":"2023-07-24T17:04:43.106797Z","iopub.status.idle":"2023-07-24T17:04:43.113776Z","shell.execute_reply.started":"2023-07-24T17:04:43.106763Z","shell.execute_reply":"2023-07-24T17:04:43.11051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# copy .yaml to the working directory\n!cp /kaggle/input/torchscript/hubmap.yaml /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2023-07-24T17:04:43.116538Z","iopub.execute_input":"2023-07-24T17:04:43.11694Z","iopub.status.idle":"2023-07-24T17:04:44.141268Z","shell.execute_reply.started":"2023-07-24T17:04:43.116887Z","shell.execute_reply":"2023-07-24T17:04:44.139725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"yolo_model = YOLO('yolov8n.pt')  # load a pretrained model (recommended for training)\n\n# Train the model with 2 GPUs\nyolo_model.train(data='/kaggle/working/hubmap.yaml', epochs=100, imgsz=IMG_SIZE, device=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"yolo_model.export()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run Inference on Images","metadata":{}},{"cell_type":"code","source":"yolo_model = YOLO('/kaggle/input/torchscript/yolo_100epoch_200im.pt')\nsam_model = SamModel.from_pretrained(\"facebook/sam-vit-base\")\nsam_model.load_state_dict(torch.load(\"/kaggle/input/torchscript/sam_20epoch_222ims.pt\"))\n\nsam_model.eval()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = HubMapDataset(\"/kaggle/input/hubmap-hacking-the-human-vasculature/train\", \"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\", processor, num_samples=300)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# One image\nidx = 68\nimg_id = train_dset.labels.iloc[idx]['id']\nprint(img_id)\nprint(train_dset.labels.iloc[idx]['type'])\n\nimg_dir = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train\"\nimg = cv2.imread(f\"{img_dir}/{img_id}.tif\")\n\nwith torch.no_grad():\n    yolo_res = yolo_model.predict(img,)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T19:24:44.120561Z","iopub.execute_input":"2023-07-24T19:24:44.121081Z","iopub.status.idle":"2023-07-24T19:24:44.186949Z","shell.execute_reply.started":"2023-07-24T19:24:44.121037Z","shell.execute_reply":"2023-07-24T19:24:44.185958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, ax = plt.subplots(1,2)\n\nbxs = train_dset[idx]['input_boxes'].squeeze().numpy().astype(np.int32)\nimg_bb = cv2.rectangle(cv2.resize(img, (1024,1024)), (bxs[0], bxs[1]), (bxs[2], bxs[3]), (0,255,0), BBOX_STROKE)\n\nres_plotted = yolo_res[0].plot()\n\nax[0].imshow(res_plotted)\nax[1].imshow(img_bb)","metadata":{"execution":{"iopub.status.busy":"2023-07-24T19:33:01.439997Z","iopub.execute_input":"2023-07-24T19:33:01.440506Z","iopub.status.idle":"2023-07-24T19:33:03.181423Z","shell.execute_reply.started":"2023-07-24T19:33:01.440463Z","shell.execute_reply":"2023-07-24T19:33:03.180439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-07-24T19:30:02.718941Z","iopub.execute_input":"2023-07-24T19:30:02.719305Z","iopub.status.idle":"2023-07-24T19:30:03.420469Z","shell.execute_reply.started":"2023-07-24T19:30:02.719276Z","shell.execute_reply":"2023-07-24T19:30:03.419391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():  \n    \n    for batch in tqdm(train_dataloader):\n        yolo_res = yolo_model.predict(batch[\"pixel_values\"])\n        # forward pass\n        outputs = model(pixel_values=batch[\"pixel_values\"].to(device),\n                      input_boxes=yolo_res.boxes.to(device),\n                      multimask_output=False)\n\n        # compute loss\n        predicted_masks = outputs.pred_masks.squeeze(1)\n        ground_truth_masks = batch[\"ground_truth_mask\"].squeeze().float().to(device)\n\n        ground_truth_masks = nn.functional.interpolate(ground_truth_masks.unsqueeze(0),\n                size=(256, 256),\n                mode='bilinear',\n                align_corners=False)\n\n        assert predicted_masks.squeeze().shape == ground_truth_masks.squeeze().shape, f\"{predicted_masks.shape} does not match {ground_truth_masks.shape}\"\n\n        loss = seg_loss(predicted_masks.squeeze(), ground_truth_masks.squeeze())\n        # print(loss)\n\n        # backward pass (compute gradients of parameters w.r.t. loss)\n        optimizer.zero_grad()\n        loss.backward()\n\n        # optimize\n        optimizer.step()\n        epoch_losses.append(loss.item())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Old","metadata":{}},{"cell_type":"code","source":"# Test xyxy to cxcywh\ni = 50\nj = 1\nout_str = \"\"\n\nlabel = df.iloc[i]['annotations'][j]\nprint(df.iloc[i]['id'])\n\n# get bounding box\nground_truth_mask = np.array(label['coordinates'])\nmask = cv2.fillPoly(np.zeros((img.shape[0], img.shape[1])), [ground_truth_mask.reshape((-1,1,2))], color=(255,0,0))\nbbox = get_bounding_box(mask)\ncxcy_bbox = xyxy_to_cxcywh(bbox, mask)\n\nline = str(label['type']) + \" \"\nline += ' '.join(map(str, cxcy_bbox))\n\nout_str += line + \"\\n\"\nout_str","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mean_and_std(dataloader):\n    channels_sum, channels_squared_sum, num_batches = 0, 0, 0\n    for data in tqdm(dataloader):\n        # Mean over batch, height and width, but not over the channels\n        channels_sum += torch.mean(data['pixel_values'], dim=[0,2,3])\n        channels_squared_sum += torch.mean(data['pixel_values']**2, dim=[0,2,3])\n        num_batches += 1\n        if num_batches > 5:\n            break\n    \n    mean = channels_sum / num_batches\n\n    # std = sqrt(E[X^2] - (E[X])^2)\n    std = (channels_squared_sum / num_batches - mean ** 2) ** 0.5\n\n    return mean, std\ntrain_dset = HubMapDataset(\"/kaggle/input/hubmap-hacking-the-human-vasculature/train\", \"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\", processor)\ntrain_dataloader = DataLoader(train_dset, batch_size=64, shuffle=True)\nmean, std = get_mean_and_std(train_dataloader)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transformer = torch.nn.Sequential(\n    transforms.Resize((N_IMG_SIZE, N_IMG_SIZE)),\n    # transforms.Normalize(tuple(mean.tolist()), tuple(std.tolist())),\n)\nscripted_transforms = torch.jit.script(transformer)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalized_img = scripted_transforms(img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(1,2)\nimg = np.transpose(next(iter(train_dataloader))['pixel_values'][0], (1,2,0))\nnormalized_img = scripted_transforms(img)\n\naxarr[0].imshow(img)\naxarr[1].imshow(normalized_img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gt_mask = train_dset[idx]['ground_truth_mask']\nprompt = get_bounding_box(gt_mask)\n\ninputs = processor(img, input_boxes=[[prompt]], return_tensors=\"pt\").to(device)\nfor k,v in inputs.items():\n    print(k,v.shape)\nprint(inputs['reshaped_input_sizes'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Modified image (resized)\nprint(train_dset.labels.iloc[idx]['type'])\nplt.imshow(img_proc)\nplt.imshow(cv2.resize(gt_mask, (img_proc.shape[0], img_proc.shape[1])), alpha=0.5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}