{"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 skimage \nfrom skimage import morphology\nimport pandas as pd, numpy as np\nimport json \nimport cv2\nimport torch \nimport matplotlib.pyplot as plt\nimport torchvision\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nimport os\nimport collections\nimport albumentations as A\nimport random\nfrom matplotlib.patches import Rectangle\nimport time\nimport sys\nimport torchvision.transforms as transforms\nUSE_CV= False\n \nif USE_CV:\n    !pip install pycocotools\n    sys.path.append(\"/kaggle/input/detection-wheel\")\n    from engine import train_one_epoch, evaluate\n    import utils\n\nif torch.cuda.is_available(): \n    dev = \"cuda:0\" \nelse: \n    dev = \"cpu\" \n    \ndevice = torch.device(dev) #tensor.to(device)\n\nTRAIN_PATH= '/kaggle/input/hubmap-hacking-the-human-vasculature/train'\nJSON_PATH='/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl'\nRESUME_TRAINING = False\nif RESUME_TRAINING:\n    model_path= '/kaggle/input/hubmap-training/6_epochs_model.pth'\n\n\nuse_ds2= True\nDILATION= True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-14T15:17:44.980280Z","iopub.execute_input":"2023-06-14T15:17:44.981264Z","iopub.status.idle":"2023-06-14T15:17:52.665476Z","shell.execute_reply.started":"2023-06-14T15:17:44.981224Z","shell.execute_reply":"2023-06-14T15:17:52.664061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"LB contains only ds1. So valid should oly have ds1. Train both on ds1 and ds2 and only on ds1 -> better or worse.\nTry dilation.","metadata":{}},{"cell_type":"code","source":"def fix_all_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n    \nfix_all_seeds(1)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:52.667647Z","iopub.execute_input":"2023-06-14T15:17:52.668039Z","iopub.status.idle":"2023-06-14T15:17:52.681220Z","shell.execute_reply.started":"2023-06-14T15:17:52.668006Z","shell.execute_reply":"2023-06-14T15:17:52.680199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bad_ids= ['18be061202ea', \n '29d2f472e46d',\n '4e455f0cb054',\n '8e90e6189c6b',\n '8f256d18b5e4',\n '90481ae2a0c9',\n 'ba276097772d',\n 'd1d485660263',\n 'd850250778f2',\n 'f45a29109ff5',\n 'f48f6580655c'] #These all have only unsure annotations","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:52.683004Z","iopub.execute_input":"2023-06-14T15:17:52.684586Z","iopub.status.idle":"2023-06-14T15:17:52.690938Z","shell.execute_reply.started":"2023-06-14T15:17:52.684548Z","shell.execute_reply":"2023-06-14T15:17:52.689585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_meta = pd.read_csv('/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv')\n#tile_meta= tile_meta[~tile_meta.id.isin(bad_ids)].reset_index()\n\nif USE_CV:\n    ds1_ids= tile_meta.loc[tile_meta.dataset==1,'id'].values.tolist()\n    ds2_ids= tile_meta.loc[tile_meta.dataset==2,'id'].values.tolist()\n    random.shuffle(ds1_ids)\n    TRAIN_SIZE= int(len(ds1_ids)*0.8)\n    train_ids= ds1_ids[:TRAIN_SIZE]\n    valid_ids = ds1_ids[TRAIN_SIZE:]\n\nelse:\n    all_ids= tile_meta.id.values.tolist()","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:52.692349Z","iopub.execute_input":"2023-06-14T15:17:52.692794Z","iopub.status.idle":"2023-06-14T15:17:52.756301Z","shell.execute_reply.started":"2023-06-14T15:17:52.692762Z","shell.execute_reply":"2023-06-14T15:17:52.754975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = [0.485, 0.456, 0.406]  # Mean values for RGB channels\nstd = [0.229, 0.224, 0.225]  # Standard deviation values for RGB channels\n\n# Define the transformation pipeline including normalization\ntransform = transforms.Compose([\n    transforms.ToTensor(),  # Convert the image to a tensor, also changes ordering of channels\n    transforms.Normalize(mean, std)  # Normalize the tensor\n    \n])","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:52.759685Z","iopub.execute_input":"2023-06-14T15:17:52.760066Z","iopub.status.idle":"2023-06-14T15:17:52.766138Z","shell.execute_reply.started":"2023-06-14T15:17:52.760032Z","shell.execute_reply":"2023-06-14T15:17:52.765090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(JSON_PATH,'r') as json_file: \n    json_labels= [json.loads(line) for line in json_file]","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:52.767233Z","iopub.execute_input":"2023-06-14T15:17:52.767581Z","iopub.status.idle":"2023-06-14T15:17:57.773995Z","shell.execute_reply.started":"2023-06-14T15:17:52.767552Z","shell.execute_reply":"2023-06-14T15:17:57.772774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubMapDataset(torch.utils.data.Dataset):\n    def __init__ (self, train_path,json_files,transform,allowed_ids, img_size=512): \n        \n        self.allowed_ids= allowed_ids \n        self.json_labels= [x for x in json_files if x['id'] in allowed_ids] # {x:y for x,y in dict.items() if x in allowed_ids}\n        self.train_path = train_path\n        self.transform= transform\n        self.img_size= img_size #512 if not working\n        #self.img_dict = collections.defaultdict(dict)\n        \n\n    def __len__(self):\n        return len(self.json_labels)\n        \n    def resize(self,img, interp):\n            return cv2.resize(img, (self.img_size, self.img_size), interpolation=interp)\n        \n    def __getitem__(self, idx): # img_id= list(dict.keys())[idx]\n            img_id = self.json_labels[idx]\n            img_path= f\"{self.train_path}/{img_id['id']}.tif\"\n            img= cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB) #return np array of pixels \n            #img = img/255 #normalize pixels (0-1)\n            \n            #Count how many objects there are in image:\n            count=0\n            for annot in img_id['annotations']:\n                if annot['type'] == 'blood_vessel' or 'glomerulus' or'unsure':\n                    count+=1\n                   \n            mask= np.zeros((count,img.shape[0], img.shape[1]), dtype=np.uint8)\n            boxes=[]\n            labels= [1 for _ in range(count)]\n            idx=0\n            \n            for annot in img_id['annotations']:\n                coords= annot['coordinates']\n                \n                if annot['type']=='blood_vessel' or'unsure':\n                    lines = np.array(coords)\n                    lines = lines.reshape(-1, 1, 2)\n                    cv2.fillPoly(mask[idx,:,:], [lines],1 ) #Fills in the mask (in place) , replace 1 with idx+1\n                    idx+=1\n                    \n                    for cord in coords:\n                        rr, cc = np.array([i[1] for i in cord]), np.asarray([i[0] for i in cord]) #row goes first (aka y axis)\n                        y_max, y_min = np.max(rr), np.min(rr)\n                        x_max, x_min = np.max(cc), np.min(cc)\n                        boxes.append([x_min,y_min,x_max,y_max])  \n                        \n                elif annot['type']=='glomerulus':\n                    lines = np.array(coords)\n                    lines = lines.reshape(-1, 1, 2)\n                    cv2.fillPoly(mask[idx,:,:], [lines], 1) #Fills in the mask (in place)\n                    labels[idx]=2\n                    idx+=1                   \n                    \n                    for cord in coords:\n                        rr, cc = np.array([i[1] for i in cord]), np.asarray([i[0] for i in cord]) #row goes first (aka y axis)\n                        y_max, y_min = np.max(rr), np.min(rr)\n                        x_max, x_min = np.max(cc), np.min(cc)\n                        boxes.append([x_min,y_min,x_max,y_max])\n                        \n            #labels = [1 for _ in range(count)] \n\n            boxes = torch.as_tensor(boxes, dtype=torch.float32)\n            labels = torch.as_tensor(labels, dtype=torch.int64)\n            \n            for i,m in enumerate(mask):\n                mask[i]= morphology.binary_dilation(m)\n                \n            if img.shape != (self.img_size, self.img_size, 3):\n                for i,m in enumerate(mask):\n                    mask[i] = self.resize(m, cv2.INTER_NEAREST)\n                img = self.resize(img,cv2.INTER_NEAREST)\n            \n            if self.transform:\n                img= transform(img)\n                \n            mask = torch.as_tensor(mask, dtype=torch.uint8)\n\n\n            area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])  #y_max - y_min, x_max - x_min\n            iscrowd = torch.zeros((count,), dtype=torch.int64)\n\n            # This is the required target for the Mask R-CNN\n            target = {\n                'boxes': boxes,\n                'labels': labels,\n                'masks': mask,\n                'image_id': torch.tensor([idx]),\n                'area': area, \n                'iscrowd': iscrowd\n            }  \n            \n        \n            return img , target\n                \n                        \n                        ","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:32:28.652212Z","iopub.execute_input":"2023-06-14T15:32:28.652608Z","iopub.status.idle":"2023-06-14T15:32:28.678391Z","shell.execute_reply.started":"2023-06-14T15:32:28.652579Z","shell.execute_reply":"2023-06-14T15:32:28.677530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mask_rcnn(num_classes):\n    \n    model = torchvision.models.detection.maskrcnn_resnet50_fpn(weights=\"DEFAULT\")\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n\n    # now get the number of input features for the mask classifier\n    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n    hidden_layer = 256\n    # and replace the mask predictor with a new one\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask,\n                                                       hidden_layer,\n                                                       num_classes)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:57.806014Z","iopub.execute_input":"2023-06-14T15:17:57.807064Z","iopub.status.idle":"2023-06-14T15:17:57.829528Z","shell.execute_reply.started":"2023-06-14T15:17:57.807020Z","shell.execute_reply":"2023-06-14T15:17:57.828099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MOMENTUM = 0.9\nLEARNING_RATE = 0.001 #from 0.001\nWEIGHT_DECAY = 0.0005\nMASK_THRESHOLD = 0.5\nUSE_SCHEDULER = False","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:17:57.831013Z","iopub.execute_input":"2023-06-14T15:17:57.832029Z","iopub.status.idle":"2023-06-14T15:17:57.843282Z","shell.execute_reply.started":"2023-06-14T15:17:57.831978Z","shell.execute_reply":"2023-06-14T15:17:57.841673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 3 #2 -> background (0),bv(1), glomerulus (2)\n\nif USE_CV:\n    ds_train= HubMapDataset(TRAIN_PATH,json_labels, True,train_ids)\n    if use_ds2:\n        ds_valid= HubMapDataset(TRAIN_PATH,json_labels, True,valid_ids + ds2_ids)\n    else: \n        ds_valid= HubMapDataset(TRAIN_PATH,json_labels, True,valid_ids)\n    dl_train = torch.utils.data.DataLoader(ds_train, batch_size=8, shuffle=True,collate_fn=lambda x: tuple(zip(*x))) \n    dl_val = torch.utils.data.DataLoader(ds_valid, batch_size=4, shuffle=False, collate_fn= lambda x: tuple(zip(*x)))\n    \nelse:\n    dataset = HubMapDataset(TRAIN_PATH,json_labels, True, all_ids)\n    dl= torch.utils.data.DataLoader(dataset, batch_size=8, shuffle=True,collate_fn=lambda x: tuple(zip(*x))) \n    \n    #next(iter(dl))","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:32:31.883484Z","iopub.execute_input":"2023-06-14T15:32:31.883918Z","iopub.status.idle":"2023-06-14T15:32:32.939791Z","shell.execute_reply.started":"2023-06-14T15:32:31.883874Z","shell.execute_reply":"2023-06-14T15:32:32.938434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking everything is working as it should:\ndef vis_ds(idx, dataset):\n    fig, ax= plt.subplots()\n    img, tgt = dataset[idx]\n    img = np.transpose(img.numpy().astype('uint8'), (1,2,0)) #CHW -> HWC for plotting\n    ax.imshow(img)\n    full_mask= np.zeros((512,512))\n    for m in tgt['masks']:\n        full_mask = np.logical_or(full_mask, m)\n    ax.imshow(full_mask, alpha=0.4)\n    box_coords= tgt['boxes']\n    for x in range(len(box_coords)):\n        rect = Rectangle((box_coords[x][0], box_coords[x][1]), box_coords[x][2]-box_coords[x][0], box_coords[x][3]-box_coords[x][1], linewidth=1, edgecolor='r', facecolor='none')\n        ax.add_patch(rect)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:32:59.206919Z","iopub.execute_input":"2023-06-14T15:32:59.207364Z","iopub.status.idle":"2023-06-14T15:32:59.218141Z","shell.execute_reply.started":"2023-06-14T15:32:59.207331Z","shell.execute_reply":"2023-06-14T15:32:59.216824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_CV:\n    vis_ds(6, ds_train)\nelse:\n    vis_ds(6, dataset)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:33:00.842336Z","iopub.execute_input":"2023-06-14T15:33:00.843173Z","iopub.status.idle":"2023-06-14T15:33:01.283559Z","shell.execute_reply.started":"2023-06-14T15:33:00.843087Z","shell.execute_reply":"2023-06-14T15:33:01.282177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model= get_mask_rcnn(3)\nif RESUME_TRAINING:\n    model.load_state_dict(torch.load(model_path)) #training from last saved weights\nNUM_EPOCHS= 8\nmodel.to(device)\nmodel.train()\noptimizer = torch.optim.SGD(model.parameters(), lr=LEARNING_RATE,momentum=MOMENTUM, weight_decay=WEIGHT_DECAY) #change to Adam\n#lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n\n#n_batches, n_batches_val = len(dl_train), len(dl_val) #205 batches in total\n\n#validation_losses = []\ndef train_model(model, dataloader, dataloader_val= None):\n\n    #min_train_loss=10000\n    n_batches= len(dataloader)\n\n    for epoch in range(1, NUM_EPOCHS + 1):\n        print(f\"Starting epoch {epoch} of {NUM_EPOCHS}\")\n\n        time_start = time.time()\n        loss_accum = 0.0\n        loss_mask_accum = 0.0\n\n        for batch_idx, (images, targets) in enumerate(dataloader, 1):\n\n            # Predict\n            images = list(image.to(device) for image in images)\n            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n            loss_dict = model(images, targets)\n            loss = sum(loss for loss in loss_dict.values())\n\n            # Backprop\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n            # Logging\n            loss_mask = loss_dict['loss_mask'].item()\n            loss_accum += loss.item()\n            loss_mask_accum += loss_mask\n\n            if batch_idx % 50 == 0:\n                print(f\"    [Batch {batch_idx:3d} / {n_batches:3d}] Batch train loss: {loss.item():7.3f}. Mask-only loss: {loss_mask:7.3f}\")\n\n        if USE_SCHEDULER:\n            lr_scheduler.step()\n\n        # Train losses\n        train_loss = loss_accum / n_batches\n        train_loss_mask = loss_mask_accum / n_batches\n\n        #Validation:\n        if dataloader_val is not None:\n            evaluate(model, dl_val, device=device)\n\n        torch.save(model.state_dict(), f\"{epoch}_epochs_model.pth\")\n        prefix = f\"[Epoch {epoch:2d} / {NUM_EPOCHS:2d}]\"\n        print(prefix)\n        print(f\"{prefix} Train mask-only loss: {train_loss_mask:7.3f}\")\n        #print(f\"{prefix} Val mask-only loss  : {val_loss_mask:7.3f}\")\n        print(prefix)\n        print(f\"{prefix} Train loss: {train_loss:7.3f}\")\n        # print(\"Val loss: {val_loss:7.3f} [{elapsed:.0f} secs]\")\n        #print(prefix)\n\n    #torch.save(model.state_dict(), f\"{NUM_EPOCHS}_epochs_model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:38:28.284629Z","iopub.execute_input":"2023-06-14T15:38:28.285087Z","iopub.status.idle":"2023-06-14T15:38:30.291793Z","shell.execute_reply.started":"2023-06-14T15:38:28.285055Z","shell.execute_reply":"2023-06-14T15:38:30.290896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(model,dl)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:38:34.805883Z","iopub.execute_input":"2023-06-14T15:38:34.806327Z","iopub.status.idle":"2023-06-14T15:38:57.798049Z","shell.execute_reply.started":"2023-06-14T15:38:34.806282Z","shell.execute_reply":"2023-06-14T15:38:57.796180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"    # Validation \n    val_loss_accum = 0\n    val_loss_mask_accum = 0 \n    \n    with torch.no_grad():\n        for batch_idx, (images, targets) in enumerate(dl_val, 1):\n            images = list(image.to(device) for image in images)\n            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n            \n            val_loss_dict = model(images, targets)\n            val_batch_loss = sum(loss for loss in val_loss_dict.values())\n            val_loss_accum += val_batch_loss.item()\n            val_loss_mask_accum += val_loss_dict['loss_mask'].item()\n    \n    # Validation losses\n    val_loss = val_loss_accum / n_batches_val\n    val_loss_mask = val_loss_mask_accum / n_batches_val\n    elapsed = time.time() - time_start\n    \n    validation_losses.append(val_loss)","metadata":{}},{"cell_type":"code","source":"# Take a few examples from val_dl\ndef show_results(idx, ds):\n    \n    img, targets = ds[idx]\n    plt.imshow(img.numpy().transpose((1,2,0)))\n    plt.title(\"Image\")\n    plt.show()\n    \n    masks = np.zeros((512, 512))\n    for mask in targets['masks']:\n        masks = np.logical_or(masks, mask)\n    plt.imshow(img.numpy().transpose((1,2,0)))\n    plt.imshow(masks, alpha=0.3)\n    plt.title(\"Ground truth\")\n    plt.show()\n    \n    model.eval()\n    with torch.no_grad():\n        preds = model([img.to(device)])[0]\n\n    plt.imshow(img.cpu().numpy().transpose((1,2,0)))\n    all_preds_masks = np.zeros((512, 512))\n    for i,mask in enumerate(preds['masks'].cpu().detach().numpy()):\n        if preds['labels'][i]==1:\n            all_preds_masks = np.logical_or(all_preds_masks, mask[0] > MASK_THRESHOLD)\n    plt.imshow(all_preds_masks, alpha=0.4)\n    plt.title(\"Predictions\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:18:00.004985Z","iopub.status.idle":"2023-06-14T15:18:00.005561Z","shell.execute_reply.started":"2023-06-14T15:18:00.005343Z","shell.execute_reply":"2023-06-14T15:18:00.005364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_results(6, dataset)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:18:00.007741Z","iopub.status.idle":"2023-06-14T15:18:00.008569Z","shell.execute_reply.started":"2023-06-14T15:18:00.008327Z","shell.execute_reply":"2023-06-14T15:18:00.008351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_results(7, dataset)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T15:18:00.010028Z","iopub.status.idle":"2023-06-14T15:18:00.010853Z","shell.execute_reply.started":"2023-06-14T15:18:00.010622Z","shell.execute_reply":"2023-06-14T15:18:00.010647Z"},"trusted":true},"execution_count":null,"outputs":[]}]}