{"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":"#### I couldn't run this notebook on the kaggle environment; needs too much GPU & RAM.\n#### Feel free to try it on your own machine (if it has sufficient hardware).","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn\nfrom matplotlib import style\nimport random\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom tqdm import tqdm\nfrom PIL import Image","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(\"../input/sartorius-cell-instance-segmentation/train.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_cells = train.groupby(by=\"id\")[\"annotation\"].agg(lambda x: list(x)).reset_index()[\"annotation\"].map(len)\nn_cells.plot(kind=\"hist\")\nplt.title(\"number of cells per image: distribution\")\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_cells.quantile(0.98)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BOX_DETECTIONS_PER_IMG = 550\nNUM_CLASSES = 2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = \"../input/sartorius-cell-instance-segmentation/train\" + \"/c4121689002f.png\"\nimg = plt.imread(path)\nplt.imshow(img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"liss = [5,8,6,4,2]\nfor i, x in enumerate(liss):\n    print(f\"i:{i}, x:{x}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arr = np.array([[1,0,0,1],\n                [0,1,0,0],\n                [1,1,1,0]])\nnp.where(arr)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DatasetObject(Dataset):\n    \n    def __init__(self, img_dir, df):\n        super().__init__()\n        self.dir = img_dir\n        self.df = df\n        self.instances = self.df.groupby(by=\"id\")[\"annotation\"].agg(lambda x: list(x)).reset_index()\n        self.height=520\n        self.width=704\n        \n    def __len__(self):\n        return len(self.instances)\n    \n    def get_img(self, img_id):\n        path = \"../input/sartorius-cell-instance-segmentation/train\" + f\"/{img_id}\" + \".png\"\n        return Image.open(path)#.convert(\"RGB\")\n    \n    def get_submask(self, annot, width=704, height=520):\n        a = annot.split()\n        submask = np.zeros(shape=(height * width), dtype=float)\n        start = list(map(int, a[0::2]))\n        n = list(map(int, a[1::2]))\n        end = []\n        for (s, n) in zip(start, n):\n            end.append(s+n-1)\n        for s, e in zip(start, end):\n            submask[s:e] = 1\n        submask = submask.reshape((height, width))\n        return submask\n    \n    def get_box(self, submask):\n        masked_pxls = np.where(submask)\n        heights, widths = masked_pxls\n        hmin, hmax, wmin, wmax = min(heights), max(heights), min(widths), max(widths)\n        return wmin, hmin, wmax, hmax\n    \n    def get_all_submasks(self, img_id, width=704, height=520):\n        #print(f\"width={width}\\nimg_id={img_id}\")\n        #print(\"self=\",self)\n        annots = list(self.df[self.df[\"id\"] == img_id][\"annotation\"])\n        n=len(annots)\n        #print(f\"width={width}\\nimg_id={img_id}\")\n        shape = (n, height, width)\n        #print(shape)\n        submasks = np.zeros(shape=shape)\n        for i, ann in enumerate(annots):\n            submasks[i,:] = self.get_submask(ann)\n        return submasks\n    \n    def get_all_boxes(self, img_id):\n        submasks = self.get_all_submasks(img_id)\n        n = submasks.shape[0]\n        boxes = np.zeros(shape=(n, 4))\n        for i in range(n):\n            boxes[i,:] = self.get_box(submasks[i])\n        return boxes\n    \n    def get_mask(self, img_id, width=704, height=520):\n        annots = list(self.df[self.df[\"id\"] == img_id][\"annotation\"])\n        mask = np.zeros(shape=(height, width), dtype=float)\n        for ann in annots:\n            get_submask(ann)\n            mask+=submask\n        return mask\n    \n    def __getitem__(self, ind):\n        #print(\"ind=\",ind)\n        img_id = self.instances.loc[ind, \"id\"]\n        img = self.get_img(img_id)\n        submasks = self.get_all_submasks(img_id=img_id)\n        boxes = self.get_all_boxes(img_id)\n        n = boxes.shape[0]\n        labels = [1 for i in range(n)]\n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n        labels = torch.as_tensor(labels, dtype=torch.int64)\n        submasks = torch.as_tensor(submasks, dtype=torch.uint8)\n        target = dict({\"boxes\":boxes, \"masks\":submasks, \"labels\":labels})\n        img = np.array(img.convert(\"RGB\"))\n        img = torch.as_tensor(img, dtype=torch.float32).view(3,self.height,self.width)\n        return img, target","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_dir = \"../input/sartorius-cell-instance-segmentation/train\"\ndataset = DatasetObject(img_dir=img_dir, df=train)\ncollater = lambda x: tuple(zip(*x))\nloader = DataLoader(dataset, batch_size=1, shuffle=True, collate_fn=collater, num_workers=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_maskrcnn(hidden_layer=256):\n    model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=True, \n                                                               box_detections_per_img=BOX_DETECTIONS_PER_IMG)\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\n    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES)\n    \n    return model\n\nDEVICE=torch.device('cuda')\nmodel = get_maskrcnn()\nmodel.to(DEVICE)\nfor param in model.parameters():\n    param.requires_grad = True\nmodel.train();","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"threshold_list = [0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9, 0.95]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou_1(mask1, mask2):\n    mask2[(mask2>0.5)]=1\n    mask2[(mask2<=0.5)]=0\n    inter = (mask1.cpu().numpy() & mask2.detach().cpu().numpy().astype(int)).sum()\n    union = (mask1.cpu().numpy() | mask2.detach().cpu().numpy().astype(int)).sum()\n    return inter/union","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou_2(mask1,mask2):\n    return iou_1(1-mask1,1-mask2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou(mask1,mask2):\n    return (iou_1(mask1,mask2) + iou_2(mask1,mask2)) / 2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou_matrix(true, pred):\n    pred = pred.squeeze()\n    ntrue = true.shape[0]\n    npred = pred.shape[0]\n    matrix = np.zeros(shape=(ntrue, npred))\n    for i in range(ntrue):\n        for j in range(npred):\n            matrix[i,j] = iou(true[i,:,:], pred[j,:,:])\n    return matrix","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou_thresh(mat, threshold):\n    matrix = mat.copy()\n    m, n = matrix.shape\n    for i in range(m):\n        row = matrix[i,:]\n        max_ind = row.argmax()\n        if (row>threshold).sum()>=1:\n            matrix[i,:] = 0\n            matrix[max_ind] = 1\n        else:\n            matrix[i,:]=0\n    return matrix","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def av_precision(mat):\n    av_prec = 0\n    for t in threshold_list:\n        av_prec += iou_thresh(mat,t)\n    av_prec /= len(threshold_list)\n    return av_prec","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_score(true, pred):\n    if pred.shape[0]==0:\n        return 0\n    else:\n        matrix = iou_matrix(true, pred)\n        return av_precision(matrix)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"params = model.parameters()\noptimizer = torch.optim.Adam(params)#,lr=0.0001,momentum=0.9,weight_decay=0.0005)\nn_epochs = 1\nnsteps = len(loader)\nfor e in range(n_epochs):\n    print(f\"epoch {e+1}\")\n    epoch_loss = 0\n    #progress_bar = tqdm(range(nsteps))\n    for batch_ind, (images, targets) in enumerate(loader,1):\n        model.train();\n        print(\"batch\",batch_ind)\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        #print(targets[0][\"masks\"].shape[0],targets[1][\"masks\"].shape[0])\n        loss_dict = model(images, targets)\n        loss = sum(loss for loss in loss_dict.values())\n        \"\"\"print(images[0].shape)\n        print(targets[0][\"boxes\"].shape)\n        print(targets[0][\"masks\"].shape)\n        print(targets[0][\"labels\"].shape)\n        print(loss)\"\"\"\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss\n        #progress_bar.update(1)\n        if (batch_ind)%10 == 0:\n            model.eval();\n            outs = model(images)\n            metric = get_score(targets[0][\"masks\"],outs[0][\"masks\"])\n            print(f\"batch: {batch_ind}, loss={loss:.2f}, metric={metric:.2f}%\")\n            del outs, metric, loss, loss_dict, images, targets\n    epoch_loss /= len(loader)\n    print(f\"Average Epoch Training Loss: {epoch_loss}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}