{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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":296547,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":253838,"modelId":275267}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport random\nimport colorsys\nfrom tqdm.auto import tqdm\n\nimport cv2\nimport albumentations as A\n\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nimport torchvision\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.transforms import functional as F\n\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:03:48.988026Z","iopub.execute_input":"2025-03-22T16:03:48.988362Z","iopub.status.idle":"2025-03-22T16:03:48.993244Z","shell.execute_reply.started":"2025-03-22T16:03:48.988336Z","shell.execute_reply":"2025-03-22T16:03:48.992198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/sartorius-cell-instance-segmentation","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:03:52.692909Z","iopub.execute_input":"2025-03-22T16:03:52.693229Z","iopub.status.idle":"2025-03-22T16:03:52.828896Z","shell.execute_reply.started":"2025-03-22T16:03:52.693206Z","shell.execute_reply":"2025-03-22T16:03:52.8278Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROOT = \"/kaggle/input/sartorius-cell-instance-segmentation\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:03:55.073324Z","iopub.execute_input":"2025-03-22T16:03:55.073651Z","iopub.status.idle":"2025-03-22T16:03:55.077791Z","shell.execute_reply.started":"2025-03-22T16:03:55.073625Z","shell.execute_reply":"2025-03-22T16:03:55.076805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(ROOT + \"/train.csv\")\ndisplay(train_df.head(2))\nprint(f'Train data shape: {train_df.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:00.446102Z","iopub.execute_input":"2025-03-22T16:04:00.446392Z","iopub.status.idle":"2025-03-22T16:04:00.794327Z","shell.execute_reply.started":"2025-03-22T16:04:00.446371Z","shell.execute_reply":"2025-03-22T16:04:00.793432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cls_map = {value:idx for idx,value in enumerate(train_df[\"cell_type\"].unique())}\ncls_map","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:05.020035Z","iopub.execute_input":"2025-03-22T16:04:05.020357Z","iopub.status.idle":"2025-03-22T16:04:05.02901Z","shell.execute_reply.started":"2025-03-22T16:04:05.02033Z","shell.execute_reply":"2025-03-22T16:04:05.028267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_targets_mask(df, img_id):\n    \"\"\"\n    Function to get target masks\n    \n    rles contains mask's description\n    \"\"\"\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    \n    return targets, rles","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:07.755312Z","iopub.execute_input":"2025-03-22T16:04:07.755614Z","iopub.status.idle":"2025-03-22T16:04:07.759801Z","shell.execute_reply.started":"2025-03-22T16:04:07.755592Z","shell.execute_reply":"2025-03-22T16:04:07.758993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_id = train_df.iloc[0][\"id\"]\nimg = cv2.imread(f'{ROOT}/train/{img_id}.png')\nplt.imshow(img)\nlabels, rles = get_targets_mask(train_df, img_id)\nlen(rles)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:11.625166Z","iopub.execute_input":"2025-03-22T16:04:11.625474Z","iopub.status.idle":"2025-03-22T16:04:11.934232Z","shell.execute_reply.started":"2025-03-22T16:04:11.625451Z","shell.execute_reply":"2025-03-22T16:04:11.933414Z"}},"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    \n    return colors","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:19.905594Z","iopub.execute_input":"2025-03-22T16:04:19.905901Z","iopub.status.idle":"2025-03-22T16:04:19.910432Z","shell.execute_reply.started":"2025-03-22T16:04:19.905877Z","shell.execute_reply":"2025-03-22T16:04:19.909549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def decode_rle_mask(rle_mask, shape=(520, 704)):\n\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    rle_mask - it is a string with start and length coordinate values\n    \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 = [np.asarray(x, dtype=int) for x in (rle_mask[0:][::2], rle_mask[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n\n    mask = np.zeros((shape[0] * shape[1]), dtype=np.uint8)\n    \n    for start, end in zip(starts, ends):\n        mask[start:end] = 1\n\n    mask = mask.reshape(shape[0], shape[1])\n    # Transform mask value (necessary step before Image augmentation)\n    mask = np.uint8(mask)\n    \n    return mask\n\n\ndef encode_rle_mask(mask, shape=(520, 704)):\n    \"\"\"\n    (Used for create submission file)\n    Parameters\n    ----------\n    mask with a given shape\n    \"\"\"\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    \n    return rle.tolist()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:22.749779Z","iopub.execute_input":"2025-03-22T16:04:22.750131Z","iopub.status.idle":"2025-03-22T16:04:22.756898Z","shell.execute_reply.started":"2025-03-22T16:04:22.750101Z","shell.execute_reply":"2025-03-22T16:04:22.75587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"masks = []\n\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    \n    masks.append(decoded_mask)\n\nlen(masks)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:28.782537Z","iopub.execute_input":"2025-03-22T16:04:28.782954Z","iopub.status.idle":"2025-03-22T16:04:28.862571Z","shell.execute_reply.started":"2025-03-22T16:04:28.782907Z","shell.execute_reply":"2025-03-22T16:04:28.861603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"masks_stack = np.stack(masks)\nmasks_stack.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:31.552658Z","iopub.execute_input":"2025-03-22T16:04:31.553025Z","iopub.status.idle":"2025-03-22T16:04:31.62706Z","shell.execute_reply.started":"2025-03-22T16:04:31.552991Z","shell.execute_reply":"2025-03-22T16:04:31.626268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_bboxes_from_mask(masks):\n    \"\"\"\n    Function to get bboxes in coco format\n    \"\"\"\n    coco_boxes = []\n    \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    \n    # Преобразуем в формат np для ДатаЛоадер  \n    coco_boxes = np.asarray(coco_boxes)\n    \n    return coco_boxes\n\nbboxes = get_bboxes_from_mask(masks_stack)\nbboxes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:37.200223Z","iopub.execute_input":"2025-03-22T16:04:37.200515Z","iopub.status.idle":"2025-03-22T16:04:38.024048Z","shell.execute_reply.started":"2025-03-22T16:04:37.200494Z","shell.execute_reply":"2025-03-22T16:04:38.023267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_mask(image, mask, color, alpha=0.5):\n    \"\"\"\n    Apply the given mask to the image\n    \"\"\"\n    for c in range(3):\n        image[:, :, c] = np.where(mask == 1,\n                                  image[:, :, c] *\n                                  (1 - alpha) + alpha * color[c] * 255,\n                                  image[:, :, c])\n    \n    return image\n\n\n\ndef plot_image_annotations(image, masks, bboxes, labels, aug=None):\n    \"\"\"\n    Function to plot masks and bboxes with augmentation\n    \"\"\"\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,\n                        labels=labels)\n        \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(image, (box[2], box[3]), (box[0], box[1]), color=color, thickness=2)\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":"2025-03-22T16:04:44.319998Z","iopub.execute_input":"2025-03-22T16:04:44.320299Z","iopub.status.idle":"2025-03-22T16:04:44.327281Z","shell.execute_reply.started":"2025-03-22T16:04:44.320276Z","shell.execute_reply":"2025-03-22T16:04:44.326422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks_stack, bboxes=bboxes, labels=labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:04:48.791404Z","iopub.execute_input":"2025-03-22T16:04:48.791889Z","iopub.status.idle":"2025-03-22T16:04:54.75977Z","shell.execute_reply.started":"2025-03-22T16:04:48.791862Z","shell.execute_reply":"2025-03-22T16:04:54.758775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_augmentations = A.Compose([\n    A.Resize(640, 640),\n    A.HorizontalFlip(),\n    A.VerticalFlip(),\n    A.RandomRotate90()\n], bbox_params={\"format\": \"pascal_voc\",\n                \"min_area\": 0,\n                \"min_visibility\": 0,\n                'label_fields': ['labels']\n               })","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:00.398209Z","iopub.execute_input":"2025-03-22T16:05:00.398613Z","iopub.status.idle":"2025-03-22T16:05:00.405816Z","shell.execute_reply.started":"2025-03-22T16:05:00.398581Z","shell.execute_reply":"2025-03-22T16:05:00.404743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks_stack, bboxes=bboxes, labels=labels, aug=train_augmentations)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:05.642988Z","iopub.execute_input":"2025-03-22T16:05:05.64332Z","iopub.status.idle":"2025-03-22T16:05:13.543504Z","shell.execute_reply.started":"2025-03-22T16:05:05.643294Z","shell.execute_reply":"2025-03-22T16:05:13.542447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CellSegData(torch.utils.data.Dataset):\n    \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        \n        if split=='train':\n            self.dataset = train\n        else:\n            self.dataset = test\n        \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        \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=img.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(image=image, masks=masks, bboxes=bboxes,\n                            labels=labels)\n            image = augmented['image']\n            masks = augmented['masks']\n            bboxes = augmented['bboxes']\n            bboxes = np.stack(bboxes).astype(int)\n        \n        # Transform to np\n        masks = np.asarray(masks)\n        \n        bboxes = torch.as_tensor(bboxes, dtype=torch.int64)\n\n\n        is_bad_labels = False\n        degenerate_boxes = bboxes[:, 2:] <= bboxes[:, :2]\n        \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        \n        return torch.Tensor(image), target #torch.Tensor(image)/255, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:20.05783Z","iopub.execute_input":"2025-03-22T16:05:20.058172Z","iopub.status.idle":"2025-03-22T16:05:20.069605Z","shell.execute_reply.started":"2025-03-22T16:05:20.058145Z","shell.execute_reply":"2025-03-22T16:05:20.068598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = CellSegData(f'{ROOT}/train', train_df, 'train', train_augmentations, cls_map)\ndataset_test = CellSegData(f'{ROOT}/train', train_df, 'test', train_augmentations, cls_map)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:26.764432Z","iopub.execute_input":"2025-03-22T16:05:26.76499Z","iopub.status.idle":"2025-03-22T16:05:30.054113Z","shell.execute_reply.started":"2025-03-22T16:05:26.764919Z","shell.execute_reply":"2025-03-22T16:05:30.053135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:32.90538Z","iopub.execute_input":"2025-03-22T16:05:32.905683Z","iopub.status.idle":"2025-03-22T16:05:32.929959Z","shell.execute_reply.started":"2025-03-22T16:05:32.90566Z","shell.execute_reply":"2025-03-22T16:05:32.929206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import datetime\nimport errno\nimport os\nimport time\nfrom collections import defaultdict, deque\n\nimport torch\nimport torch.distributed as dist\n\n\nclass SmoothedValue:\n    \"\"\"Track a series of values and provide access to smoothed values over a\n    window or the global series average.\n    \"\"\"\n\n    def __init__(self, window_size=20, fmt=None):\n        if fmt is None:\n            fmt = \"{median:.4f} ({global_avg:.4f})\"\n        self.deque = deque(maxlen=window_size)\n        self.total = 0.0\n        self.count = 0\n        self.fmt = fmt\n\n    def update(self, value, n=1):\n        self.deque.append(value)\n        self.count += n\n        self.total += value * n\n\n    def synchronize_between_processes(self):\n        \"\"\"\n        Warning: does not synchronize the deque!\n        \"\"\"\n        if not is_dist_avail_and_initialized():\n            return\n        t = torch.tensor([self.count, self.total], dtype=torch.float64, device=\"cuda\")\n        dist.barrier()\n        dist.all_reduce(t)\n        t = t.tolist()\n        self.count = int(t[0])\n        self.total = t[1]\n\n    @property\n    def median(self):\n        d = torch.tensor(list(self.deque))\n        return d.median().item()\n\n    @property\n    def avg(self):\n        d = torch.tensor(list(self.deque), dtype=torch.float32)\n        return d.mean().item()\n\n    @property\n    def global_avg(self):\n        return self.total / self.count\n\n    @property\n    def max(self):\n        return max(self.deque)\n\n    @property\n    def value(self):\n        return self.deque[-1]\n\n    def __str__(self):\n        return self.fmt.format(\n            median=self.median, avg=self.avg, global_avg=self.global_avg, max=self.max, value=self.value\n        )\n\n\ndef all_gather(data):\n    \"\"\"\n    Run all_gather on arbitrary picklable data (not necessarily tensors)\n    Args:\n        data: any picklable object\n    Returns:\n        list[data]: list of data gathered from each rank\n    \"\"\"\n    world_size = get_world_size()\n    if world_size == 1:\n        return [data]\n    data_list = [None] * world_size\n    dist.all_gather_object(data_list, data)\n    return data_list\n\n\ndef reduce_dict(input_dict, average=True):\n    \"\"\"\n    Args:\n        input_dict (dict): all the values will be reduced\n        average (bool): whether to do average or sum\n    Reduce the values in the dictionary from all processes so that all processes\n    have the averaged results. Returns a dict with the same fields as\n    input_dict, after reduction.\n    \"\"\"\n    world_size = get_world_size()\n    if world_size < 2:\n        return input_dict\n    with torch.inference_mode():\n        names = []\n        values = []\n        # sort the keys so that they are consistent across processes\n        for k in sorted(input_dict.keys()):\n            names.append(k)\n            values.append(input_dict[k])\n        values = torch.stack(values, dim=0)\n        dist.all_reduce(values)\n        if average:\n            values /= world_size\n        reduced_dict = {k: v for k, v in zip(names, values)}\n    return reduced_dict\n\n\nclass MetricLogger:\n    def __init__(self, delimiter=\"\\t\"):\n        self.meters = defaultdict(SmoothedValue)\n        self.delimiter = delimiter\n\n    def update(self, **kwargs):\n        for k, v in kwargs.items():\n            if isinstance(v, torch.Tensor):\n                v = v.item()\n            assert isinstance(v, (float, int))\n            self.meters[k].update(v)\n\n    def __getattr__(self, attr):\n        if attr in self.meters:\n            return self.meters[attr]\n        if attr in self.__dict__:\n            return self.__dict__[attr]\n        raise AttributeError(f\"'{type(self).__name__}' object has no attribute '{attr}'\")\n\n    def __str__(self):\n        loss_str = []\n        for name, meter in self.meters.items():\n            loss_str.append(f\"{name}: {str(meter)}\")\n        return self.delimiter.join(loss_str)\n\n    def synchronize_between_processes(self):\n        for meter in self.meters.values():\n            meter.synchronize_between_processes()\n\n    def add_meter(self, name, meter):\n        self.meters[name] = meter\n\n    def log_every(self, iterable, print_freq, header=None):\n        i = 0\n        if not header:\n            header = \"\"\n        start_time = time.time()\n        end = time.time()\n        iter_time = SmoothedValue(fmt=\"{avg:.4f}\")\n        data_time = SmoothedValue(fmt=\"{avg:.4f}\")\n        space_fmt = \":\" + str(len(str(len(iterable)))) + \"d\"\n        if torch.cuda.is_available():\n            log_msg = self.delimiter.join(\n                [\n                    header,\n                    \"[{0\" + space_fmt + \"}/{1}]\",\n                    \"eta: {eta}\",\n                    \"{meters}\",\n                    \"time: {time}\",\n                    \"data: {data}\",\n                    \"max mem: {memory:.0f}\",\n                ]\n            )\n        else:\n            log_msg = self.delimiter.join(\n                [header, \"[{0\" + space_fmt + \"}/{1}]\", \"eta: {eta}\", \"{meters}\", \"time: {time}\", \"data: {data}\"]\n            )\n        MB = 1024.0 * 1024.0\n        for obj in iterable:\n            data_time.update(time.time() - end)\n            yield obj\n            iter_time.update(time.time() - end)\n            if i % print_freq == 0 or i == len(iterable) - 1:\n                eta_seconds = iter_time.global_avg * (len(iterable) - i)\n                eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))\n                if torch.cuda.is_available():\n                    print(\n                        log_msg.format(\n                            i,\n                            len(iterable),\n                            eta=eta_string,\n                            meters=str(self),\n                            time=str(iter_time),\n                            data=str(data_time),\n                            memory=torch.cuda.max_memory_allocated() / MB,\n                        )\n                    )\n                else:\n                    print(\n                        log_msg.format(\n                            i, len(iterable), eta=eta_string, meters=str(self), time=str(iter_time), data=str(data_time)\n                        )\n                    )\n            i += 1\n            end = time.time()\n        total_time = time.time() - start_time\n        total_time_str = str(datetime.timedelta(seconds=int(total_time)))\n        print(f\"{header} Total time: {total_time_str} ({total_time / len(iterable):.4f} s / it)\")\n\n\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n\n\ndef mkdir(path):\n    try:\n        os.makedirs(path)\n    except OSError as e:\n        if e.errno != errno.EEXIST:\n            raise\n\n\ndef setup_for_distributed(is_master):\n    \"\"\"\n    This function disables printing when not in master process\n    \"\"\"\n    import builtins as __builtin__\n\n    builtin_print = __builtin__.print\n\n    def print(*args, **kwargs):\n        force = kwargs.pop(\"force\", False)\n        if is_master or force:\n            builtin_print(*args, **kwargs)\n\n    __builtin__.print = print\n\n\ndef is_dist_avail_and_initialized():\n    if not dist.is_available():\n        return False\n    if not dist.is_initialized():\n        return False\n    return True\n\n\ndef get_world_size():\n    if not is_dist_avail_and_initialized():\n        return 1\n    return dist.get_world_size()\n\n\ndef get_rank():\n    if not is_dist_avail_and_initialized():\n        return 0\n    return dist.get_rank()\n\n\ndef is_main_process():\n    return get_rank() == 0\n\n\ndef save_on_master(*args, **kwargs):\n    if is_main_process():\n        torch.save(*args, **kwargs)\n\n\ndef init_distributed_mode(args):\n    if \"RANK\" in os.environ and \"WORLD_SIZE\" in os.environ:\n        args.rank = int(os.environ[\"RANK\"])\n        args.world_size = int(os.environ[\"WORLD_SIZE\"])\n        args.gpu = int(os.environ[\"LOCAL_RANK\"])\n    elif \"SLURM_PROCID\" in os.environ:\n        args.rank = int(os.environ[\"SLURM_PROCID\"])\n        args.gpu = args.rank % torch.cuda.device_count()\n    else:\n        print(\"Not using distributed mode\")\n        args.distributed = False\n        return\n\n    args.distributed = True\n\n    torch.cuda.set_device(args.gpu)\n    args.dist_backend = \"nccl\"\n    print(f\"| distributed init (rank {args.rank}): {args.dist_url}\", flush=True)\n    torch.distributed.init_process_group(\n        backend=args.dist_backend, init_method=args.dist_url, world_size=args.world_size, rank=args.rank\n    )\n    torch.distributed.barrier()\n    setup_for_distributed(args.rank == 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:38.795703Z","iopub.execute_input":"2025-03-22T16:05:38.796046Z","iopub.status.idle":"2025-03-22T16:05:38.817885Z","shell.execute_reply.started":"2025-03-22T16:05:38.796018Z","shell.execute_reply":"2025-03-22T16:05:38.817025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pycocotools","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:47.246139Z","iopub.execute_input":"2025-03-22T16:05:47.246438Z","iopub.status.idle":"2025-03-22T16:05:50.538376Z","shell.execute_reply.started":"2025-03-22T16:05:47.246416Z","shell.execute_reply":"2025-03-22T16:05:50.537185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_loader = torch.utils.data.DataLoader(\n    dataset, batch_size=2, shuffle=True, num_workers=2, prefetch_factor=2, collate_fn=collate_fn)\n\ndata_loader_test = torch.utils.data.DataLoader(\n    dataset_test, batch_size=1, shuffle=False, num_workers=4, collate_fn=collate_fn) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:54.320414Z","iopub.execute_input":"2025-03-22T16:05:54.32073Z","iopub.status.idle":"2025-03-22T16:05:54.324919Z","shell.execute_reply.started":"2025-03-22T16:05:54.320707Z","shell.execute_reply":"2025-03-22T16:05:54.32406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir -p '/root/.cache/torch/hub/checkpoints'\n!cp ../input/maskrcnn_resnet50_fpn/pytorch/default/1/resnet50-0676ba61.pth /root/.cache/torch/hub/checkpoints/resnet50-0676ba61.pth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:05:58.730832Z","iopub.execute_input":"2025-03-22T16:05:58.731189Z","iopub.status.idle":"2025-03-22T16:05:59.197768Z","shell.execute_reply.started":"2025-03-22T16:05:58.731159Z","shell.execute_reply":"2025-03-22T16:05:59.196803Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_CLASSES = 2\n\nmodel = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=False, box_detections_per_img=600)\n\n# 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\nhidden_layer = 256\n# and replace the mask predictor with a new one\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES)\n# weights = torch.load('weights/checkpoint.pth')\n# move model to the right device\n# model.state_dict(weights['model'])\nmodel.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:06:05.380642Z","iopub.execute_input":"2025-03-22T16:06:05.380993Z","iopub.status.idle":"2025-03-22T16:06:06.236886Z","shell.execute_reply.started":"2025-03-22T16:06:05.380965Z","shell.execute_reply":"2025-03-22T16:06:06.236036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.AdamW(params, lr=0.0001)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,\n                                               step_size=3,\n                                               gamma=0.1)\n\noutput_dir = 'weights'\nos.makedirs(output_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:06:18.755664Z","iopub.execute_input":"2025-03-22T16:06:18.756034Z","iopub.status.idle":"2025-03-22T16:06:18.761806Z","shell.execute_reply.started":"2025-03-22T16:06:18.756006Z","shell.execute_reply":"2025-03-22T16:06:18.760959Z"}},"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    \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        \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    \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        \n        all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n        \n    plt.imshow(all_preds_masks, alpha=0.4)\n    plt.title(\"Predictions\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:06:22.881878Z","iopub.execute_input":"2025-03-22T16:06:22.882244Z","iopub.status.idle":"2025-03-22T16:06:22.889395Z","shell.execute_reply.started":"2025-03-22T16:06:22.882215Z","shell.execute_reply":"2025-03-22T16:06:22.88835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\n\nfor epoch in tqdm(range(num_epochs)):\n    # train for one epoch, printing every 10 iterations\n    model.train()\n        \n    metric_logger = MetricLogger(delimiter=\"  \")\n    metric_logger.add_meter(\"lr\", SmoothedValue(window_size=1, fmt=\"{value:.6f}\"))\n    header = f\"Epoch: [{epoch}]\"\n\n    lr_scheduler = None\n        \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        losses = sum(loss for loss in loss_dict.values())\n\n        # reduce losses over all GPUs for logging purposes\n        loss_dict_reduced = 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    # evaluate on the test dataset\n    # evaluate(model, data_loader_test, device=device)\n    if output_dir:\n        checkpoint = {\n            \"model\": model.state_dict(),\n            \"optimizer\": optimizer.state_dict(),\n            \"epoch\": epoch,\n        }\n        save_on_master(checkpoint, os.path.join(output_dir, f\"model_{epoch}.pth\"))\n        save_on_master(checkpoint, os.path.join(output_dir, \"checkpoint.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:06:52.331797Z","iopub.execute_input":"2025-03-22T16:06:52.332146Z","iopub.status.idle":"2025-03-22T16:36:57.879228Z","shell.execute_reply.started":"2025-03-22T16:06:52.332116Z","shell.execute_reply":"2025-03-22T16:36:57.877709Z"}},"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        image = Image.open(image_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image, _ = self.transforms(image=image, target=None)\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":"2025-03-22T16:41:11.88592Z","iopub.execute_input":"2025-03-22T16:41:11.886265Z","iopub.status.idle":"2025-03-22T16:41:11.891846Z","shell.execute_reply.started":"2025-03-22T16:41:11.886238Z","shell.execute_reply":"2025-03-22T16:41:11.890907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nclass Compose:\n    def __init__(self, transforms):\n        self.transforms = transforms\n\n    def __call__(self, image, target):\n        for t in self.transforms:\n            image, target = t(image, target)\n        return image, target\n    \n\nclass Normalize:\n    def __call__(self, image, target):\n        image = F.normalize(image, RESNET_MEAN, RESNET_STD)\n        return image, target\n    \n    \nclass ToTensor:\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        \n        return image, target\n    \n    \ndef get_transform(train):\n    transforms = [ToTensor()]\n    if True:\n        transforms.append(Normalize())\n\n    return Compose(transforms)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:41:18.211145Z","iopub.execute_input":"2025-03-22T16:41:18.211466Z","iopub.status.idle":"2025-03-22T16:41:18.217276Z","shell.execute_reply.started":"2025-03-22T16:41:18.211438Z","shell.execute_reply":"2025-03-22T16:41:18.216289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds_test = CellTestDataset(f'{ROOT}/test', transforms=get_transform(train=False))\nds_test[2]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:41:22.195373Z","iopub.execute_input":"2025-03-22T16:41:22.195692Z","iopub.status.idle":"2025-03-22T16:41:22.235393Z","shell.execute_reply.started":"2025-03-22T16:41:22.195664Z","shell.execute_reply":"2025-03-22T16:41:22.234726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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))\n\n\ndef 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:41:26.098076Z","iopub.execute_input":"2025-03-22T16:41:26.098408Z","iopub.status.idle":"2025-03-22T16:41:26.103695Z","shell.execute_reply.started":"2025-03-22T16:41:26.09838Z","shell.execute_reply":"2025-03-22T16:41:26.102872Z"}},"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    \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        \n    submission.append((image_id, rle))\n    print(submission)\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":"2025-03-22T16:41:31.589686Z","iopub.execute_input":"2025-03-22T16:41:31.590022Z","iopub.status.idle":"2025-03-22T16:43:58.156695Z","shell.execute_reply.started":"2025-03-22T16:41:31.589996Z","shell.execute_reply":"2025-03-22T16:43:58.15588Z"}},"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":"2025-03-22T16:44:12.474178Z","iopub.execute_input":"2025-03-22T16:44:12.47449Z","iopub.status.idle":"2025-03-22T16:44:12.489637Z","shell.execute_reply.started":"2025-03-22T16:44:12.474466Z","shell.execute_reply":"2025-03-22T16:44:12.488781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/working","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T16:44:17.726273Z","iopub.execute_input":"2025-03-22T16:44:17.726619Z","iopub.status.idle":"2025-03-22T16:44:17.915544Z","shell.execute_reply.started":"2025-03-22T16:44:17.726591Z","shell.execute_reply":"2025-03-22T16:44:17.914699Z"}},"outputs":[],"execution_count":null}]}