{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14763465,"sourceType":"datasetVersion","datasetId":9436567},{"sourceId":14764072,"sourceType":"datasetVersion","datasetId":9436830},{"sourceId":14771147,"sourceType":"datasetVersion","datasetId":9441731}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Install all .whl files using pip\n! pip install --no-index --find-links /kaggle/input/pip-deps albumentations\n! pip install --no-index --find-links /kaggle/input/pip-deps pandas\n! pip install --no-index --find-links /kaggle/input/pip-deps opencv-python\n! pip install --no-index --find-links /kaggle/input/pip-deps torchvision\n! pip install --no-index --find-links /kaggle/input/pip-deps tqdm\n! pip install --no-index --find-links /kaggle/input/pip-deps scikit-learn\n! pip install --no-index --find-links /kaggle/input/pip-deps matplotlib","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:24:20.410960Z","iopub.execute_input":"2026-02-08T15:24:20.411316Z","iopub.status.idle":"2026-02-08T15:24:42.857461Z","shell.execute_reply.started":"2026-02-08T15:24:20.411288Z","shell.execute_reply":"2026-02-08T15:24:42.856582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\n\n\nimport pandas as pd\nimport numpy as np\n\nimport colorsys\nimport cv2\nfrom PIL import Image\nimport albumentations as A\nfrom tqdm.auto import tqdm\n\nimport torch\n\nfrom sklearn.model_selection import train_test_split\n\nimport matplotlib.pyplot as plt\n\nimport torchvision\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:24:42.859154Z","iopub.execute_input":"2026-02-08T15:24:42.859455Z","iopub.status.idle":"2026-02-08T15:25:24.328050Z","shell.execute_reply.started":"2026-02-08T15:24:42.859419Z","shell.execute_reply":"2026-02-08T15:25:24.327497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n\nsys.path.append('/kaggle/input/pytorch-utils/references/detection')\nprint(os.listdir('/kaggle/input/pytorch-utils/references/detection'))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.328929Z","iopub.execute_input":"2026-02-08T15:25:24.329399Z","iopub.status.idle":"2026-02-08T15:25:24.349790Z","shell.execute_reply.started":"2026-02-08T15:25:24.329366Z","shell.execute_reply":"2026-02-08T15:25:24.349222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import utils","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.351415Z","iopub.execute_input":"2026-02-08T15:25:24.351814Z","iopub.status.idle":"2026-02-08T15:25:24.372770Z","shell.execute_reply.started":"2026-02-08T15:25:24.351792Z","shell.execute_reply":"2026-02-08T15:25:24.372271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_colors(N, bright=True):\n    \"\"\"\n    Generate random colors.\n    To get visually distinct colors, generate them in HSV space then\n    convert to RGB.\n    \"\"\"\n    brightness = 1.0 if bright else 0.7\n    hsv = [(i / N, 1, brightness) for i in range(N)]\n    colors = list(map(lambda c: colorsys.hsv_to_rgb(*c), hsv))\n    random.shuffle(colors)\n    return colors","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.373472Z","iopub.execute_input":"2026-02-08T15:25:24.373724Z","iopub.status.idle":"2026-02-08T15:25:24.378568Z","shell.execute_reply.started":"2026-02-08T15:25:24.373703Z","shell.execute_reply":"2026-02-08T15:25:24.377844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def decode_rle_mask(rle_mask, shape=(520, 704)):\n    \"\"\"\n    Decode run-length encoded segmentation mask string into 2d array\n\n    Parameters\n    ----------\n    rle_mask (str): Run-length encoded segmentation mask string\n    shape (tuple): Height and width of the mask\n\n    Returns\n    -------\n    mask [numpy.ndarray of shape (height, width)]: Decoded 2d segmentation mask\n    \"\"\"\n\n    rle_mask = rle_mask.split()\n    starts, lengths = [\n        np.asarray(x, dtype=int) for x in (rle_mask[0:][::2], rle_mask[1:][::2])\n    ]\n    starts -= 1\n    ends = starts + lengths\n\n    mask = np.zeros((shape[0] * shape[1]), dtype=np.uint8)\n    for start, end in zip(starts, ends):\n        mask[start:end] = 1\n\n    mask = mask.reshape(shape[0], shape[1])\n    mask = np.uint8(mask)\n    return mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:39:52.684691Z","iopub.execute_input":"2026-02-08T15:39:52.685009Z","iopub.status.idle":"2026-02-08T15:39:52.691134Z","shell.execute_reply.started":"2026-02-08T15:39:52.684983Z","shell.execute_reply":"2026-02-08T15:39:52.690161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def encode_rle_mask(mask):\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    rle = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    rle[1::2] -= rle[::2]\n    return rle.tolist()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.391727Z","iopub.execute_input":"2026-02-08T15:25:24.391992Z","iopub.status.idle":"2026-02-08T15:25:24.401401Z","shell.execute_reply.started":"2026-02-08T15:25:24.391972Z","shell.execute_reply":"2026-02-08T15:25:24.400796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_bboxes_from_mask(masks):\n    coco_boxes = []\n    for mask in masks:\n        pos = np.nonzero(mask)\n        xmin = np.min(pos[1])\n        xmax = np.max(pos[1])\n        ymin = np.min(pos[0])\n        ymax = np.max(pos[0])\n        coco_boxes.append([xmin, ymin, xmax, ymax])\n    coco_boxes = np.asarray(coco_boxes)\n    return coco_boxes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.402134Z","iopub.execute_input":"2026-02-08T15:25:24.402399Z","iopub.status.idle":"2026-02-08T15:25:24.414108Z","shell.execute_reply.started":"2026-02-08T15:25:24.402369Z","shell.execute_reply":"2026-02-08T15:25:24.413511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_targets_mask(df, img_id):\n    targets = df[df[\"id\"] == img_id][\"cell_type\"].apply(lambda x: cls_map[x]).values\n    rles = df[df[\"id\"] == img_id][\"annotation\"].values\n    return targets, rles","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.414904Z","iopub.execute_input":"2026-02-08T15:25:24.415143Z","iopub.status.idle":"2026-02-08T15:25:24.426673Z","shell.execute_reply.started":"2026-02-08T15:25:24.415124Z","shell.execute_reply":"2026-02-08T15:25:24.425992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_mask(image, mask, color, alpha=0.5):\n    \"\"\"Apply the given mask to the image.\"\"\"\n    for c in range(3):\n        image[:, :, c] = np.where(\n            mask == 1,\n            image[:, :, c] * (1 - alpha) + alpha * color[c] * 255,\n            image[:, :, c],\n        )\n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.428906Z","iopub.execute_input":"2026-02-08T15:25:24.429154Z","iopub.status.idle":"2026-02-08T15:25:24.436709Z","shell.execute_reply.started":"2026-02-08T15:25:24.429135Z","shell.execute_reply":"2026-02-08T15:25:24.436034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/sartorius-cell-instance-segmentation/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.437558Z","iopub.execute_input":"2026-02-08T15:25:24.438144Z","iopub.status.idle":"2026-02-08T15:25:24.892904Z","shell.execute_reply.started":"2026-02-08T15:25:24.438122Z","shell.execute_reply":"2026-02-08T15:25:24.892269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cls_map = {value: idx for idx, value in enumerate(train_df[\"cell_type\"].unique())}\ncls_map_reversed = {cls_map[key]: key for key in cls_map.keys()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.893771Z","iopub.execute_input":"2026-02-08T15:25:24.894007Z","iopub.status.idle":"2026-02-08T15:25:24.904502Z","shell.execute_reply.started":"2026-02-08T15:25:24.893985Z","shell.execute_reply":"2026-02-08T15:25:24.903810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_id = train_df.iloc[0]['id']\nimg = cv2.imread(f'/kaggle/input/sartorius-cell-instance-segmentation/train/{img_id}.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.905490Z","iopub.execute_input":"2026-02-08T15:25:24.905785Z","iopub.status.idle":"2026-02-08T15:25:24.955108Z","shell.execute_reply.started":"2026-02-08T15:25:24.905754Z","shell.execute_reply":"2026-02-08T15:25:24.954571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(img)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:24.956051Z","iopub.execute_input":"2026-02-08T15:25:24.956354Z","iopub.status.idle":"2026-02-08T15:25:25.235542Z","shell.execute_reply.started":"2026-02-08T15:25:24.956323Z","shell.execute_reply":"2026-02-08T15:25:25.234805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels, rles = get_targets_mask(train_df, img_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.236716Z","iopub.execute_input":"2026-02-08T15:25:25.237000Z","iopub.status.idle":"2026-02-08T15:25:25.256146Z","shell.execute_reply.started":"2026-02-08T15:25:25.236977Z","shell.execute_reply":"2026-02-08T15:25:25.255342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(rles)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.257200Z","iopub.execute_input":"2026-02-08T15:25:25.257469Z","iopub.status.idle":"2026-02-08T15:25:25.267850Z","shell.execute_reply.started":"2026-02-08T15:25:25.257445Z","shell.execute_reply":"2026-02-08T15:25:25.267117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = decode_rle_mask(rles[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:40:45.162350Z","iopub.execute_input":"2026-02-08T15:40:45.162964Z","iopub.status.idle":"2026-02-08T15:40:45.199605Z","shell.execute_reply.started":"2026-02-08T15:40:45.162938Z","shell.execute_reply":"2026-02-08T15:40:45.199014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:40:46.901970Z","iopub.execute_input":"2026-02-08T15:40:46.902580Z","iopub.status.idle":"2026-02-08T15:40:47.085999Z","shell.execute_reply.started":"2026-02-08T15:40:46.902551Z","shell.execute_reply":"2026-02-08T15:40:47.085233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = decode_rle_mask(str(encode_rle_mask(mask)).replace('[', '').replace(']', '').replace(',', ''))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.518345Z","iopub.execute_input":"2026-02-08T15:25:25.518620Z","iopub.status.idle":"2026-02-08T15:25:25.523183Z","shell.execute_reply.started":"2026-02-08T15:25:25.518590Z","shell.execute_reply":"2026-02-08T15:25:25.522533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.524110Z","iopub.execute_input":"2026-02-08T15:25:25.524445Z","iopub.status.idle":"2026-02-08T15:25:25.763103Z","shell.execute_reply.started":"2026-02-08T15:25:25.524398Z","shell.execute_reply":"2026-02-08T15:25:25.762365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"masks = []\nfor mask in train_df.loc[train_df['id'] == img_id, 'annotation'].values:\n    decoded_mask = decode_rle_mask(rle_mask=mask, shape=img.shape)\n    masks.append(decoded_mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.764017Z","iopub.execute_input":"2026-02-08T15:25:25.764230Z","iopub.status.idle":"2026-02-08T15:25:25.829404Z","shell.execute_reply.started":"2026-02-08T15:25:25.764209Z","shell.execute_reply":"2026-02-08T15:25:25.828796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bboxes = get_bboxes_from_mask(masks)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:25.830397Z","iopub.execute_input":"2026-02-08T15:25:25.831086Z","iopub.status.idle":"2026-02-08T15:25:26.391418Z","shell.execute_reply.started":"2026-02-08T15:25:25.831055Z","shell.execute_reply":"2026-02-08T15:25:26.390790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_image_annotations(image, masks, bboxes, labels, aug=None):\n    image = image.copy()\n\n    colors = {unique_lbl: random_colors(1, True)[0] for unique_lbl in np.unique(labels)}\n\n    if aug is not None:\n        augmented = aug(image=image, masks=masks, bboxes=bboxes, labels=labels)\n        image = augmented[\"image\"]\n        masks = augmented[\"masks\"]\n        bboxes = augmented[\"bboxes\"]\n\n    bboxes = np.stack(bboxes).astype(int)\n\n    for idx, box in enumerate(bboxes):\n        color = tuple([int(value * 255) for value in colors[labels[idx]]])\n        image = cv2.rectangle(\n            image, (box[2], box[3]), (box[0], box[1]), color=color, thickness=2\n        )\n\n    for idx, mask in enumerate(masks):\n        color = colors[labels[idx]]\n        image = apply_mask(image, mask, color)\n\n    plt.figure(figsize=(15, 15))\n    plt.imshow(image)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:26.392417Z","iopub.execute_input":"2026-02-08T15:25:26.392714Z","iopub.status.idle":"2026-02-08T15:25:26.398935Z","shell.execute_reply.started":"2026-02-08T15:25:26.392679Z","shell.execute_reply":"2026-02-08T15:25:26.398330Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks, bboxes=bboxes, labels=labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:26.399771Z","iopub.execute_input":"2026-02-08T15:25:26.400021Z","iopub.status.idle":"2026-02-08T15:25:28.611815Z","shell.execute_reply.started":"2026-02-08T15:25:26.399997Z","shell.execute_reply":"2026-02-08T15:25:28.611059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_augmentations = A.Compose(\n    [\n        A.Resize(640, 640),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), \n    ],\n    bbox_params={\n        \"format\": \"pascal_voc\",\n        \"min_area\": 0,\n        \"min_visibility\": 0,\n        \"label_fields\": [\"labels\"],\n    },\n)\n\ntest_augmentations = A.Compose([\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), \n    A.ToTensorV2(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:28.612802Z","iopub.execute_input":"2026-02-08T15:25:28.613139Z","iopub.status.idle":"2026-02-08T15:25:28.622800Z","shell.execute_reply.started":"2026-02-08T15:25:28.613116Z","shell.execute_reply":"2026-02-08T15:25:28.622033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks, bboxes=bboxes, labels=labels, aug=train_augmentations)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:28.623939Z","iopub.execute_input":"2026-02-08T15:25:28.624189Z","iopub.status.idle":"2026-02-08T15:25:31.662422Z","shell.execute_reply.started":"2026-02-08T15:25:28.624167Z","shell.execute_reply":"2026-02-08T15:25:31.661422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CellSegData(torch.utils.data.Dataset):\n    def __init__(self, root, df, split=\"train\", aug=None, cls_map=None):\n        self.augmentations = aug\n        self.cls_map = cls_map\n\n        train, test = train_test_split(df[\"id\"].unique(), train_size=0.9)\n        if split == \"train\":\n            self.dataset = train\n        else:\n            self.dataset = test\n\n        self.dict_df = {img_id: df[df[\"id\"] == img_id] for img_id in tqdm(self.dataset)}\n        self.root = root\n\n    def __len__(self):\n        return len(self.dataset)\n\n    def __getitem__(self, index):\n        img_id = self.dataset[index]\n        image = cv2.imread(os.path.join(self.root, img_id + \".png\"))\n\n        info = self.dict_df[img_id]\n        n_objects = len(info[\"annotation\"])\n\n        labels = info[\"cell_type\"].apply(lambda x: self.cls_map[x]).values\n        rles = info[\"annotation\"].values\n\n        masks = []\n        for mask in rles:\n            decoded_mask = decode_rle_mask(rle_mask=mask, shape=image.shape)\n            masks.append(decoded_mask)\n\n        bboxes = get_bboxes_from_mask(masks)\n\n        if self.augmentations is not None:\n            augmented = self.augmentations(\n                image=image, masks=masks, bboxes=bboxes, labels=labels\n            )\n            image = augmented[\"image\"]\n            masks = augmented[\"masks\"]\n            bboxes = augmented[\"bboxes\"]\n            bboxes = np.stack(bboxes).astype(int)\n\n        masks = np.asarray(masks)\n\n        bboxes = torch.as_tensor(bboxes, dtype=torch.int64)\n\n        is_bad_labels = False\n        degenerate_boxes = bboxes[:, 2:] <= bboxes[:, :2]\n        if degenerate_boxes.any():\n            is_bad_labels = True\n            # print the first degenerate box\n            bb_idxs = torch.where(degenerate_boxes.any(dim=1))[0].numpy()\n\n        labels = torch.as_tensor([1 for i in range(len(bboxes))], dtype=torch.int64)\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n\n        if is_bad_labels:\n            bboxes = bboxes[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            labels = labels[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            masks = masks[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n\n        image_id = torch.tensor([index])\n        area = (bboxes[:, 3] - bboxes[:, 1]) * (bboxes[:, 2] - bboxes[:, 0])\n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n\n        # This is the required target for the Mask R-CNN\n        target = {\n            \"boxes\": bboxes,\n            \"labels\": labels,\n            \"masks\": masks,\n            \"image_id\": image_id,\n            \"area\": area,\n            \"iscrowd\": iscrowd,\n        }\n\n        image = image.transpose((2, 0, 1))\n        return torch.Tensor(image), target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:31.663434Z","iopub.execute_input":"2026-02-08T15:25:31.663674Z","iopub.status.idle":"2026-02-08T15:25:31.679100Z","shell.execute_reply.started":"2026-02-08T15:25:31.663653Z","shell.execute_reply":"2026-02-08T15:25:31.678342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_train_sample(model, ds_train, sample_index):\n    img, targets = ds_train[sample_index]\n    plt.imshow(img.numpy().astype(np.uint8).transpose((1, 2, 0)))\n    plt.title(\"Image\")\n    plt.show()\n\n    masks = np.zeros((640, 640))\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.cuda()])[0]\n\n    plt.imshow(img.cpu().numpy().transpose((1, 2, 0)))\n    all_preds_masks = np.zeros((640, 640))\n    for mask in preds[\"masks\"].cpu().detach().numpy():\n        all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n    plt.imshow(all_preds_masks, alpha=0.4)\n    plt.title(\"Predictions\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:31.682835Z","iopub.execute_input":"2026-02-08T15:25:31.683065Z","iopub.status.idle":"2026-02-08T15:25:31.694607Z","shell.execute_reply.started":"2026-02-08T15:25:31.683044Z","shell.execute_reply":"2026-02-08T15:25:31.694021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n\n# our dataset has two classes only - background and person\n# use our dataset and defined transformations\ndataset = CellSegData(\"/kaggle/input/sartorius-cell-instance-segmentation/train\", train_df, \"train\", train_augmentations, cls_map)\ndataset_test = CellSegData(\n    \"/kaggle/input/sartorius-cell-instance-segmentation/test\", train_df, \"test\", train_augmentations, cls_map\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:31.695500Z","iopub.execute_input":"2026-02-08T15:25:31.696002Z","iopub.status.idle":"2026-02-08T15:25:35.517455Z","shell.execute_reply.started":"2026-02-08T15:25:31.695973Z","shell.execute_reply":"2026-02-08T15:25:35.516727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define training and validation data loaders\ndata_loader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=2,\n    shuffle=True,\n    num_workers=2,\n    prefetch_factor=2,\n    collate_fn=utils.collate_fn,\n)\n\ndata_loader_test = torch.utils.data.DataLoader(\n    dataset_test,\n    batch_size=1,\n    shuffle=False,\n    num_workers=4,\n    collate_fn=utils.collate_fn,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:35.518331Z","iopub.execute_input":"2026-02-08T15:25:35.518578Z","iopub.status.idle":"2026-02-08T15:25:35.523255Z","shell.execute_reply.started":"2026-02-08T15:25:35.518556Z","shell.execute_reply":"2026-02-08T15:25:35.522724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_CLASSES = 2\n\nweights_path = '/kaggle/input/resnet-weight/resnet50-0676ba61.pth'\nresnet_state_dict = torch.load(weights_path, map_location='cpu')\n\nmodel = torchvision.models.detection.maskrcnn_resnet50_fpn(\n    weights=None, weights_backbone=None, box_detections_per_img=600\n)\n\nbackbone = model.backbone.body\nbackbone.load_state_dict(resnet_state_dict, strict=False)\nprint(\"Loaded ResNet50 weights into backbone\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:35.524169Z","iopub.execute_input":"2026-02-08T15:25:35.524547Z","iopub.status.idle":"2026-02-08T15:25:38.075829Z","shell.execute_reply.started":"2026-02-08T15:25:35.524522Z","shell.execute_reply":"2026-02-08T15:25:38.075102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# get the number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n    # replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\n\n    # now get the number of input features for the mask classifier\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:38.076769Z","iopub.execute_input":"2026-02-08T15:25:38.077048Z","iopub.status.idle":"2026-02-08T15:25:38.081968Z","shell.execute_reply.started":"2026-02-08T15:25:38.077025Z","shell.execute_reply":"2026-02-08T15:25:38.081265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"hidden_layer = 256\n    # and replace the mask predictor with a new one\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(\n    in_features_mask, hidden_layer, NUM_CLASSES\n)\n\nmodel.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:38.082840Z","iopub.execute_input":"2026-02-08T15:25:38.083108Z","iopub.status.idle":"2026-02-08T15:25:38.303674Z","shell.execute_reply.started":"2026-02-08T15:25:38.083079Z","shell.execute_reply":"2026-02-08T15:25:38.302949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    # construct an optimizer\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.AdamW(params, lr=0.0001)\n    # optimizer.load_state_dict(weights['optimizer'])\n    # and a learning rate scheduler\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:25:38.304678Z","iopub.execute_input":"2026-02-08T15:25:38.305065Z","iopub.status.idle":"2026-02-08T15:25:38.309698Z","shell.execute_reply.started":"2026-02-08T15:25:38.305032Z","shell.execute_reply":"2026-02-08T15:25:38.309119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 20\n\noutput_dir = \"weights\"\nos.makedirs(output_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:29:30.829973Z","iopub.execute_input":"2026-02-08T15:29:30.830540Z","iopub.status.idle":"2026-02-08T15:29:30.834742Z","shell.execute_reply.started":"2026-02-08T15:29:30.830503Z","shell.execute_reply":"2026-02-08T15:29:30.833962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in tqdm(range(num_epochs)):\n        # train for one epoch, printing every 10 iterations\n    model.train()\n    metric_logger = utils.MetricLogger(delimiter=\"  \")\n    metric_logger.add_meter(\n        \"lr\", utils.SmoothedValue(window_size=1, fmt=\"{value:.6f}\")\n    )\n    header = f\"Epoch: [{epoch}]\"\n\n    lr_scheduler = None\n    if epoch == 0:\n        warmup_factor = 1.0 / 1000\n        warmup_iters = min(1000, len(data_loader) - 1)\n\n        lr_scheduler = torch.optim.lr_scheduler.LinearLR(\n            optimizer, start_factor=warmup_factor, total_iters=warmup_iters\n        )\n\n    for images, targets in metric_logger.log_every(data_loader, 10, header):\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\n        losses = sum(loss for loss in loss_dict.values())\n\n            # reduce losses over all GPUs for logging purposes\n        loss_dict_reduced = utils.reduce_dict(loss_dict)\n        losses_reduced = sum(loss for loss in loss_dict_reduced.values())\n\n        loss_value = losses_reduced.item()\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        if lr_scheduler is not None:\n            lr_scheduler.step()\n\n        metric_logger.update(loss=losses_reduced, **loss_dict_reduced)\n        metric_logger.update(lr=optimizer.param_groups[0][\"lr\"])\n\n    analyze_train_sample(model, dataset, 20)\n    \n    if output_dir:\n        checkpoint = {\n            \"model\": model.state_dict(),\n            \"optimizer\": optimizer.state_dict(),\n            \"epoch\": epoch,\n        }\n        utils.save_on_master(\n            checkpoint, os.path.join(output_dir, f\"model_{epoch}.pth\")\n        )\n        utils.save_on_master(checkpoint, os.path.join(output_dir, \"checkpoint.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:29:32.601477Z","iopub.execute_input":"2026-02-08T15:29:32.601779Z","iopub.status.idle":"2026-02-08T15:37:27.023176Z","shell.execute_reply.started":"2026-02-08T15:29:32.601753Z","shell.execute_reply":"2026-02-08T15:37:27.022068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CellTestDataset(torch.utils.data.Dataset):\n    def __init__(self, image_dir, transforms=None):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.image_ids = [f[:-4]for f in os.listdir(self.image_dir)]\n        \n    \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image_path = os.path.join(self.image_dir, image_id + \".png\")\n        pil_image = Image.open(image_path).convert(\"RGB\")\n        image = np.array(pil_image)\n\n        if self.transforms is not None:\n            augmented = self.transforms(image=image)\n            image = augmented['image']\n            \n        return {'image': image, 'image_id': image_id}\n\n    \n    def __len__(self):\n        return len(self.image_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:37:39.435000Z","iopub.execute_input":"2026-02-08T15:37:39.435343Z","iopub.status.idle":"2026-02-08T15:37:39.441803Z","shell.execute_reply.started":"2026-02-08T15:37:39.435302Z","shell.execute_reply":"2026-02-08T15:37:39.441208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds_test = CellTestDataset('/kaggle/input/sartorius-cell-instance-segmentation/test', transforms=test_augmentations)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:38:16.390812Z","iopub.execute_input":"2026-02-08T15:38:16.391549Z","iopub.status.idle":"2026-02-08T15:38:16.395429Z","shell.execute_reply.started":"2026-02-08T15:38:16.391520Z","shell.execute_reply":"2026-02-08T15:38:16.394723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_overlapping_pixels(mask, other_masks):\n    for other_mask in other_masks:\n        if np.sum(np.logical_and(mask, other_mask)) > 0:\n            mask[np.logical_and(mask, other_mask)] = 0\n    return mask\n\ndef rle_encoding(x):\n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\n    \n    for b in dots:\n        if (b>prev+1): run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n        \n    return ' '.join(map(str, run_lengths))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:38:17.678241Z","iopub.execute_input":"2026-02-08T15:38:17.678586Z","iopub.status.idle":"2026-02-08T15:38:17.683825Z","shell.execute_reply.started":"2026-02-08T15:38:17.678560Z","shell.execute_reply":"2026-02-08T15:38:17.683186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval();\n\nsubmission = []\n\nfor sample in ds_test:\n    img = sample['image']\n    image_id = sample['image_id']\n    \n    with torch.no_grad():\n        result = model([img.to(device)])[0]\n    \n    previous_masks = []\n    image_rles = []\n\n    for i, mask in enumerate(result[\"masks\"]):\n        # Filter-out low-scoring results. Not tried yet.\n        score = result[\"scores\"][i].cpu().item()\n        if score < 0.3:\n            continue\n        \n        mask = mask.cpu().numpy()\n        # Keep only highly likely pixels\n        binary_mask = mask > 0.5\n        binary_mask = remove_overlapping_pixels(binary_mask, previous_masks)\n        previous_masks.append(binary_mask)\n        rle = rle_encoding(binary_mask)\n        image_rles.append(rle)\n\n        \n    submission.append((image_id, rle))\n    print((image_id, rle))\n    \n    plt.figure(figsize=(12,12))\n    ax1 = plt.subplot(121)\n    ax1.imshow(img.numpy().transpose((1,2,0)))\n    all_preds_masks = np.zeros((520, 704))\n    \n    for mask in result['masks'].cpu().detach().numpy():\n        all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n        \n    ax2 = plt.subplot(122)\n    ax2.imshow(img.numpy().transpose((1,2,0)))\n    ax2.imshow(all_preds_masks, alpha=0.3)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:41:44.671274Z","iopub.execute_input":"2026-02-08T15:41:44.672029Z","iopub.status.idle":"2026-02-08T15:41:53.021683Z","shell.execute_reply.started":"2026-02-08T15:41:44.672002Z","shell.execute_reply":"2026-02-08T15:41:53.020866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sub = pd.DataFrame(submission, columns=['id', 'predicted'])\ndf_sub.to_csv(\"submission.csv\", index=False)\ndf_sub.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:41:28.107067Z","iopub.execute_input":"2026-02-08T15:41:28.107646Z","iopub.status.idle":"2026-02-08T15:41:28.118492Z","shell.execute_reply.started":"2026-02-08T15:41:28.107616Z","shell.execute_reply":"2026-02-08T15:41:28.117738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/working","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T15:29:27.307586Z","iopub.status.idle":"2026-02-08T15:29:27.307930Z","shell.execute_reply.started":"2026-02-08T15:29:27.307756Z","shell.execute_reply":"2026-02-08T15:29:27.307778Z"}},"outputs":[],"execution_count":null}]}