{"metadata":{"kernelspec":{"display_name":"my-fenv","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.16"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12015885,"sourceType":"datasetVersion","datasetId":7559624}],"dockerImageVersionId":31041,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport time\nimport random\nimport collections\n\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport torch\nfrom torchvision import transforms\nimport torchvision\nfrom torchvision.transforms import ToPILImage\nfrom torchvision.transforms import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.models.detection import MaskRCNN_ResNet50_FPN_V2_Weights, maskrcnn_resnet50_fpn_v2, roi_heads","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Mount the drive","metadata":{}},{"cell_type":"code","source":"# #file path in google drive\n# #create a shortcut in your Drive.\n# drive.mount('/content/drive')\n# base_path = \"/content/drive/MyDrive/Sartorius_project/\"\n# data_path = os.path.join(base_path, 'sartorius-cell-instance-segmentation')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# BASE_PATH = '/content/drive/MyDrive/NYCU/Sartorius_project'\n# BASE_PATH = \"./sartorius-cell-instance-segmentation\"\nBASE_PATH = \"/kaggle/input/sartorius-cell-instance-segmentation\"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Path to zip file\n# zip_path = os.path.join(base_path, 'sartorius-cell-instance-segmentation.zip')\n# extract_to = base_path\n\n# # Extract only if not already extracted\n# if not os.path.exists(data_path) or len(os.listdir(data_path)) == 0:\n#     print(\"Extracting data...\")\n#     os.makedirs(data_path, exist_ok=True)\n#     with zipfile.ZipFile(zip_path, 'r') as zip_ref:\n#         zip_ref.extractall(extract_to)\n#     print(\"Extraction completed to:\", extract_to)\n# else:\n#     print(\"Data already extracted at:\", data_path)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nFunction to make results reproducible\n\"\"\"\ndef __set__seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed) #seed for cpu\n    torch.cuda.manual_seed(seed) #seed for gpu\n    torch.cuda.manual_seed_all(seed)\n    \n__set__seeds(2024)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_PATH = f\"{BASE_PATH}/train\"\nTEST_PATH = f\"{BASE_PATH}/test\"\nTRAIN_CSV = f\"{BASE_PATH}/train.csv\"\nTEST = False","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CSV - for test read only 5000 rows\ndf = pd.read_csv(TRAIN_CSV, nrows=5000 if TEST else None)\ndf.head()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Historgram for class distribution","metadata":{}},{"cell_type":"code","source":"df.head()\ncell_type_counts = df['cell_type'].value_counts()\nprint(cell_type_counts)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Class histogram by cell type\nplt.figure(figsize=(8, 4))\ndf['cell_type'].value_counts().plot(kind='bar', color='skyblue')\nplt.title(\"Cell Type Distribution\")\nplt.xlabel(\"Cell Type\")\nplt.ylabel(\"Number of Masks\")\nplt.xticks(rotation=45)\nplt.grid(axis='y')\nplt.tight_layout()\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Make the train/validation split from the train folder 80/20","metadata":{}},{"cell_type":"code","source":"unique_ids = df['id'].unique()\nprint(f\"Total unique images: {len(unique_ids)}\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# 90% train, 10% validation\ntrain_ids, val_ids = train_test_split(\n    unique_ids,\n    test_size=0.1,\n    random_state=42,\n    shuffle=True\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = df[df['id'].isin(train_ids)].reset_index(drop=True)\nval_df = df[df['id'].isin(val_ids)].reset_index(drop=True)\n\nprint(f\"Train images: {train_df['id'].nunique()} | Val images: {val_df['id'].nunique()}\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.to_csv(\"train_split.csv\", index=False)\nval_df.to_csv(\"val_split.csv\", index=False)\n\ntrain_image_ids = train_df['id'].unique().tolist()\nval_image_ids = val_df['id'].unique().tolist()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Visualize image and mask","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport os\nfrom PIL import Image\n\ndef rle_decode(mask_rle, shape=(520, 704)):\n    \"\"\"Decode RLE (Run-Length Encoding) encoded masks.\"\"\"\n    s = mask_rle.strip().split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0::2], s[1::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)\n\ndef visualize_sample(image_id, df, img_dir):\n    img_path = os.path.join(img_dir, f\"{image_id}.png\")\n    image = np.array(Image.open(img_path))\n\n    masks = df[df['id'] == image_id]['annotation'].tolist()\n    combined_mask = np.zeros_like(image, dtype=np.uint8)\n\n    for i, rle in enumerate(masks):\n        mask = rle_decode(rle)\n        combined_mask += mask.astype(np.uint8)\n\n    plt.figure(figsize=(12, 6))\n    plt.subplot(1, 2, 1)\n    plt.imshow(image, cmap='gray')\n    plt.title('Image')\n\n    plt.subplot(1, 2, 2)\n    plt.imshow(image, cmap='gray')\n    plt.imshow(combined_mask, alpha=0.5, cmap='jet')\n    plt.title('Image + Masks')\n    plt.show()\n\n# Example on id 0\nvisualize_sample(train_image_ids[0], df, TRAIN_PATH)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n\nIMG_WIDTH = 704\nIMG_HEIGHT = 520\n\n# To normalize\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nBATCH_SIZE = 2\n\nMOMENTUM = 0.9\nLEARNING_RATE = 1e-3\nWEIGHT_DECAY = 5e-4\n\nMASK_THRESHOLD = 0.5\n\n# Normalize to resnet mean and std\nNORMALIZE = True \n\nUSE_SCHEDULER = True\n\nNUM_EPOCHS = 50\n\nBOX_DETECTIONS_PER_IMG = 539\n\n# Dictionaries to classify each type of cell\nCELL_TYPE_DICT = {\"astro\": 1, \"cort\": 2, \"shsy5y\": 3}\nDICT_TO_CELL = {1: \"astro\", 2: \"cort\", 3: \"shsy5y\"}\nMASK_THRESHOLD_DICT = {1: 0.55, 2: 0.75, 3:  0.6}\nMIN_SCORE_DICT = {1: 0.55, 2: 0.75, 3: 0.5}\nNUM_CLASSES = len(CELL_TYPE_DICT)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Run-Length Encoding (RLE) encoding and decoding functions for masks","metadata":{}},{"cell_type":"code","source":"# def rle_decode_for_train(mask_rle, shape=(520, 704)):\n#     if pd.isnull(mask_rle):\n#         return np.zeros(shape, dtype=np.uint8)\n#     s = mask_rle.strip().split()\n#     starts, lengths = [np.asarray(x, dtype=int) for x in (s[0::2], s[1::2])]\n#     starts -= 1\n#     ends = starts + lengths\n#     img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n#     for lo, hi in zip(starts, ends):\n#         img[lo:hi] = 1\n#     return img.reshape(shape)\n\n# def rle_encode(mask):\n#     pixels = mask.T.flatten()  # column line\n#     pixels = np.concatenate([[0], pixels, [0]])\n#     runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n#     runs[1::2] -= runs[::2]\n#     return ' '.join(str(x) for x in runs)\n\ndef rle_decode(mask_rle, shape, color=1):\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.float32)\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n    return img.reshape(shape)\n\ndef rle_encoding(x):\n    '''\n    x : image to be encoded \n    Returns string of encoded image\n    '''\n    \n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\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    return ' '.join(map(str, run_lengths))","metadata":{},"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 draw_box(box):\n    result = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n    xmin = int(box[0])\n    ymin = int(box[1])\n    xmax = int(box[2])\n    ymax = int(box[3]) \n    \n    for x in range(xmin, xmax):\n        if (xmin != 0) and (xmax != IMG_WIDTH):\n            result[ymin-1][x] = 1\n            result[ymax-1][x] = 1\n            \n    for y in range(ymin, ymax):\n        if (ymin != 0) and (ymax != IMG_HEIGHT):\n            result[y][xmax-1] = 1\n            result[y][xmin-1] = 1\n            \n    return result\n\n# Get bbox of given mask\ndef get_box(a_mask):\n    pos = np.where(a_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\n    return [xmin, ymin, xmax, ymax]\n\n# Mask + img\ndef combine_masks(masks, mask_threshold):\n    maskimg = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n    for m, mask in enumerate(masks,1):\n        maskimg[mask>mask_threshold] = m\n    return maskimg\n\n# Mask + bbox\ndef combine_masks_boxes(masks, boxes):\n    result = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n    cur_max = 0\n    for i in range(IMG_WIDTH):\n        for j in range(IMG_HEIGHT):\n            \n            result[j][i] = 0\n            \n            if masks[j][i] != 0:\n                result[j][i] = masks[j][i]\n                if masks[j][i] > cur_max:\n                    cur_max = masks[j][i]\n\n    for i in range(IMG_WIDTH):\n        for j in range(IMG_HEIGHT):\n            if boxes[j][i] != 0:\n                result[j][i] = cur_max                \n    return result\n\n\"\"\"\nFilter masks using MIN_SCORE for mask \nand MAX_THRESHOLD for pixels\n\"\"\"\ndef get_filtered_masks(pred):\n    use_masks = []   \n    for i, mask in enumerate(pred[\"masks\"]):\n        scr = pred[\"scores\"][i].cpu().item()\n        label = pred[\"labels\"][i].cpu().item()\n        if scr > MIN_SCORE_DICT[label]:\n            mask = mask.cpu().numpy().squeeze()\n            # Keep only highly likely pixels\n            binary_mask = mask > MASK_THRESHOLD_DICT[label]\n            binary_mask = remove_overlapping_pixels(binary_mask, use_masks)\n            use_masks.append(binary_mask)\n\n    return use_masks","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nComputes the IoU for instance labels and predictions\n\"\"\"\ndef compute_iou(labels, y_pred, verbose=0):\n    true_objects = len(np.unique(labels))\n    pred_objects = len(np.unique(y_pred))\n\n    if verbose:\n        print(\"Number of true objects: {}\".format(true_objects))\n        print(\"Number of predicted objects: {}\".format(pred_objects))\n\n    # Compute intersection between all objects\n    intersection = np.histogram2d(\n        labels.flatten(), y_pred.flatten(), bins=(true_objects, pred_objects)\n    )[0]\n\n    # Compute areas (needed for finding the union between all objects)\n    area_true = np.histogram(labels, bins=true_objects)[0]\n    area_pred = np.histogram(y_pred, bins=pred_objects)[0]\n    area_true = np.expand_dims(area_true, -1)\n    area_pred = np.expand_dims(area_pred, 0)\n\n    # Compute union\n    union = area_true + area_pred - intersection\n    # exclude background\n    intersection = intersection[1:, 1:]\n    union = union[1:, 1:]\n    union[union == 0] = 1e-9\n    iou = intersection / union\n    \n    return iou  \n\n\"\"\"\nComputes the precision at a given threshold.\n\"\"\"\ndef precision_at(threshold, iou):\n    matches = iou > threshold\n    true_positives = np.sum(matches, axis=1) == 1  # Correct objects\n    false_positives = np.sum(matches, axis=0) == 0  # Missed objects\n    false_negatives = np.sum(matches, axis=1) == 0  # Extra objects\n    tp, fp, fn = (\n        np.sum(true_positives),\n        np.sum(false_positives),\n        np.sum(false_negatives),\n    )\n    return tp, fp, fn\n\n\"\"\"\nComputes the metric for the competition.\n\"\"\"\ndef iou_map(truths, preds, verbose=0):\n    ious = [compute_iou(truth, pred, verbose) for truth, pred in zip(truths, preds)]\n\n    if verbose:\n        print(\"Thresh\\tTP\\tFP\\tFN\\tPrec.\")\n\n    prec = []\n    for t in np.arange(0.5, 1.0, 0.05):\n        tps, fps, fns = 0, 0, 0\n        for iou in ious:\n            tp, fp, fn = precision_at(t, iou)\n            tps += tp\n            fps += fp\n            fns += fn\n\n        p = tps / (tps + fps + fns)\n        prec.append(p)\n\n        if verbose:\n            print(\"{:1.3f}\\t{}\\t{}\\t{}\\t{:1.3f}\".format(t, tps, fps, fns, p))\n\n    if verbose:\n        print(\"AP\\t-\\t-\\t-\\t{:1.3f}\".format(np.mean(prec)))\n\n    return np.mean(prec)\n\n\"\"\"\nGet average IOU mAP score for a dataset\n\"\"\"\ndef get_score(ds, mdl):\n    mdl.eval()\n    iouscore = 0\n    for i in tqdm(range(len(ds))):\n        img, targets = ds[i]\n        with torch.no_grad():\n            result = mdl([img.to(DEVICE)])[0]\n            \n        masks = combine_masks(targets['masks'], 0.5)\n        labels = pd.Series(result['labels'].cpu().numpy()).value_counts()\n\n        mask_threshold = MASK_THRESHOLD_DICT[labels.sort_values().index[-1]]\n        pred_masks = combine_masks(get_filtered_masks(result), mask_threshold)\n        iouscore += iou_map([masks],[pred_masks])\n    return iouscore / len(ds)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Transform\n# import torchvision.transforms as T\n# transform = T.Compose([\n#     T.Resize((512, 512)),\n#     T.ToTensor(),\n#     T.Normalize(\n#       mean=[0.485, 0.456, 0.406],\n#       std=[0.229, 0.224, 0.225])\n# ])\n\n\"\"\"\nAdditionally to vertical/horizontal flip (CUDA_OUT):\n-> random rotation - cells can appear any orientation\n-> gaussian blue - noise\n\"\"\"\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# class RandomRotate90:\n#     def __init__(self, prob=0.5):\n#         self.prob = prob\n\n#     def __call__(self, image, target):\n#         if random.random() < self.prob:\n#             # 90, 180, 270 degrees\n#             k = random.choice([1, 2, 3])\n#             # Rotate tensor image\n#             image = image.rot90(k, [1, 2])\n#             # Rotate masks\n#             target[\"masks\"] = target[\"masks\"].rot90(k, [1, 2])\n\n#             # Recalculate bboxes\n#             h, w = image.shape[-2:]\n#             boxes = []\n#             for mask in target[\"masks\"]:\n#                 boxes.append(get_box(mask.numpy()))\n#             target[\"boxes\"] = torch.as_tensor(boxes, dtype=torch.float32)\n#         return image, target\n\n\n# class GaussianBlur:\n#     def __init__(self, prob=0.3, kernel_size=5):\n#         self.prob = prob\n#         self.kernel_size = kernel_size\n\n#     def __call__(self, image, target):\n#         if random.random() < self.prob:\n#             image = F.gaussian_blur(image, kernel_size=self.kernel_size)\n#         return image, target\n\nclass VerticalFlip:\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-2)\n            bbox = target[\"boxes\"]\n            bbox[:, [1, 3]] = height - bbox[:, [3, 1]]\n            target[\"boxes\"] = bbox\n            target[\"masks\"] = target[\"masks\"].flip(-2)\n        return image, target\n\nclass HorizontalFlip:\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-1)\n            bbox = target[\"boxes\"]\n            bbox[:, [0, 2]] = width - bbox[:, [2, 0]]\n            target[\"boxes\"] = bbox\n            target[\"masks\"] = target[\"masks\"].flip(-1)\n        return image, target\n\nclass Normalize:\n    def __call__(self, image, target):\n        image = F.normalize(image, RESNET_MEAN, RESNET_STD)\n        return image, target\n\nclass ToTensor:\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        return image, target\n    \n    \ndef get_transform(train):\n    transforms = [ToTensor()]\n    if NORMALIZE:\n        transforms.append(Normalize())\n    \n    if train: \n        if NORMALIZE: \n            transforms.append(Normalize())\n        transforms.append(HorizontalFlip(0.5))\n        transforms.append(VerticalFlip(0.5))\n        # transforms.append(RandomRotate90(0.5))\n        # transforms.append(GaussianBlur(0.3))\n        \n    return Compose(transforms)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, image_dir, df, transforms=None):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.df = df\n        self.height = IMG_HEIGHT\n        self.width = IMG_WIDTH\n        self.image_info = collections.defaultdict(dict)\n        \n        temp_df = self.df.groupby(['id', 'cell_type'])['annotation'].agg(lambda x: list(x)).reset_index()\n        \n        for index, row in temp_df.iterrows():\n            self.image_info[index] = {\n                    'image_id': row['id'],\n                    'image_path': os.path.join(self.image_dir, row['id'] + '.png'),\n                    'annotations': row[\"annotation\"],\n                    'cell_type': CELL_TYPE_DICT[row[\"cell_type\"]]\n                    }\n\n    def __getitem__(self, idx):\n        img_path = self.image_info[idx][\"image_path\"]\n        img = Image.open(img_path).convert(\"RGB\")\n        info = self.image_info[idx]\n        n_objects = len(info['annotations'])\n        masks = np.zeros((len(info['annotations']), self.height, self.width), dtype=np.uint8)\n        boxes = []\n        \n        for i, annotation in enumerate(info['annotations']):\n            \n            a_mask = rle_decode(annotation, (IMG_HEIGHT, IMG_WIDTH))\n            a_mask = Image.fromarray(a_mask)\n            a_mask = np.array(a_mask) > 0\n            masks[i, :, :] = a_mask\n            boxes.append(get_box(a_mask))\n\n        labels = [info[\"cell_type\"] for _ in range(n_objects)]\n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n        labels = torch.as_tensor(labels, dtype=torch.int64)\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n        image_id = torch.tensor([idx])\n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n\n        target = {\n            'boxes': boxes,\n            'labels': labels,\n            'masks': masks,\n            'image_id': image_id,\n            'area': area,\n            'iscrowd': iscrowd\n        }\n\n        if self.transforms is not None:\n            img, target = self.transforms(img, target)\n        return img, target\n    \n    def __len__(self):\n        return len(self.image_info)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DataLoader for training\ntrain_dataset = CellDataset(TRAIN_PATH, df, transforms=get_transform(train=True))\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, \n                      num_workers=2, collate_fn=collate_fn)\n\n# DataLoader for validation\nval_dataset = CellDataset(TRAIN_PATH, val_df, transforms=get_transform(train=False))\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, \n                        num_workers=2, collate_fn=collate_fn)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_batch(images, targets, max_images=4):\n    plt.figure(figsize=(16, 4 * max_images))\n\n    for i in range(min(max_images, len(images))):\n        image = images[i].permute(1, 2, 0).cpu().numpy()\n        image = np.clip(image * 0.229 + 0.485, 0, 1)  # denomarlization ImageNet\n\n        masks = targets[i][\"masks\"].cpu().numpy()\n        combined_mask = np.zeros(image.shape[:2], dtype=np.uint8)\n        for mask in masks:\n            combined_mask = np.maximum(combined_mask, mask)\n\n        plt.subplot(max_images, 2, 2 * i + 1)\n        plt.imshow(image)\n        plt.title(\"Image\")\n        plt.axis('off')\n\n        plt.subplot(max_images, 2, 2 * i + 2)\n        plt.imshow(image)\n        plt.imshow(combined_mask, alpha=0.5, cmap='viridis')\n        plt.title(\"Image + Masks\")\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Récupère un batch et affiche\nimages, targets = next(iter(train_loader))\nshow_batch(images, targets)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def sample_images(ncols):\n#     '''\n#     ncols : number of images to sample\n#     Returns plot with 2x(ncols) images from training dataset \n#     '''\n#     ns = random.sample(range(100), ncols)\n#     fig, axs = plt.subplots(2, ncols, figsize=(20, 6)) \n    \n#     for i in range(ncols):\n#         img, targets = train_dataset[ns[i]]\n#         masks = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n#         boxes = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n#         axs[0][i].set_title(f\"Image {ns[i]}\")\n#         axs[0][i].imshow(img.numpy().transpose((1,2,0)))\n#         axs[0][i].axis(\"off\")\n        \n\n#         for mask in targets['masks']:\n#             box = get_box(mask)\n#             boxes = np.logical_or(boxes, draw_box(box))\n            \n#         masks = combine_masks(targets['masks'], 0.5)\n            \n#         axs[1][i].set_title(f\"{ns[i]} {DICT_TO_CELL[((targets['labels'])[0]).item()]} mask\")\n#         detections = combine_masks_boxes(masks, boxes)\n#         axs[1][i].imshow(detections)\n#         axs[1][i].axis(\"off\")\n#     plt.show()\n# sample_images(4)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"markdown","source":"https://download.pytorch.org/models/maskrcnn_resnet50_fpn_coco-bf2d0c1e.pth\nPretrained model","metadata":{}},{"cell_type":"code","source":"# Load COCO-pretrained weights\nmodel = torchvision.models.detection.maskrcnn_resnet50_fpn(weights=None, weights_backbone=None)\nstate_dict = torch.load(\"//kaggle/input/pretrained/maskrcnn_resnet50_fpn_coco_0.15.1.pth\")\nmodel.load_state_dict(state_dict)\n\n# Replace classification head\nin_features = model.roi_heads.box_predictor.cls_score.in_features\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES + 1)\n\n# Replace mask head\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\nhidden_layer = 128\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES + 1)\n\nmodel.to(DEVICE)\n\nfor param in model.parameters():\n    param.requires_grad = True\n    \nmodel.train()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self, patience=10, delta=0.001):\n        self.patience = patience\n        self.delta = delta\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.best_epoch = 0\n\n    def __call__(self, val_loss, epoch):\n        if self.best_score is None:\n            self.best_score = val_loss\n            self.best_epoch = epoch\n        elif val_loss > self.best_score - self.delta:\n            self.counter += 1\n            print(f\"EarlyStopping: {self.counter}/{self.patience} without improvement.\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = val_loss\n            self.best_epoch = epoch\n            self.counter = 0","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\n\noptimizer = torch.optim.SGD(params, lr=LEARNING_RATE, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\nn_batches, n_batches_val = len(train_loader), len(val_loader)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"validation_mask_losses = []\ntrain_losses = []\nval_losses = []\n\nfor epoch in range(1, NUM_EPOCHS + 1):\n    time_start = time.time()\n    loss_accum = 0.0\n    loss_mask_accum = 0.0\n    loss_classifier_accum = 0.0\n    \n    for batch_idx, (images, targets) in enumerate(tqdm(train_loader), 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        loss_dict = model(images, targets)\n        loss = sum(loss for loss in loss_dict.values())\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        loss_mask = loss_dict['loss_mask'].item()\n        loss_accum += loss.item()\n        loss_mask_accum += loss_mask\n        loss_classifier_accum += loss_dict['loss_classifier'].item()\n        \n    # Train losses\n    train_loss = loss_accum / n_batches\n    train_loss_mask = loss_mask_accum / n_batches\n    train_loss_classifier = loss_classifier_accum / n_batches\n\n    if USE_SCHEDULER and epoch >= 5:\n        lr_scheduler.step()\n    \n    # Validation\n    # model.eval()\n    val_loss_epoch = 0\n    val_loss_mask_accum = 0\n    val_loss_classifier_accum = 0\n    \n    with torch.no_grad():\n        for batch_idx, (images, targets) in enumerate(tqdm(val_loader), 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_epoch += val_batch_loss.item()\n            val_loss_mask_accum += val_loss_dict['loss_mask'].item()\n            val_loss_classifier_accum += val_loss_dict['loss_classifier'].item()\n\n    val_loss = val_loss_epoch / n_batches_val\n    val_loss_mask = val_loss_mask_accum / n_batches_val\n    val_loss_classifier = val_loss_classifier_accum / n_batches_val\n    #time per epoch\n    epoch_time = time.time() - time_start\n    \n    #for plotting\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n    validation_mask_losses.append(val_loss_mask)\n    \n    torch.save(model.state_dict(), f\"pytorch_model-e{epoch}.bin\")\n    print(f\"[Epoch {epoch} / {NUM_EPOCHS}] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n    print(f\"[Epoch {epoch} / {NUM_EPOCHS}] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n    print(f\"[Epoch {epoch} / {NUM_EPOCHS}] Train loss: {train_loss:7.3f}. Val loss: {val_loss:7.3f}\")\n    print(f\"Time for epoch {epoch}: {epoch_time:.2f} seconds\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#training vall loss curve\nplt.plot(train_losses, label=\"Train Loss\")\nplt.plot(val_losses, label=\"Val Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.title(\"Training & Validation Loss Curve\")\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plots: the image, The image + the ground truth mask + detection boxes, The image + the predicted mask + predicted boxes\ndef analyze_sample(model, train_dataset, sample_index):\n    img, targets = train_dataset[sample_index]\n    fig, axs = plt.subplots(1, 3, figsize=(20, 40), facecolor=\"#fefefe\") \n    \n    masks = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n    boxes = np.zeros((IMG_HEIGHT, IMG_WIDTH))\n    \n    axs[0].imshow(img.numpy().transpose((1,2,0)))\n    axs[0].set_title(\"Image\")\n    axs[0].axis(\"off\")\n    \n    for mask in targets['masks']:\n        box = get_box(mask)\n        boxes = np.logical_or(boxes, draw_box(box)) \n        \n    masks = combine_masks(targets['masks'], 0.5)\n    detections = combine_masks_boxes(masks, boxes)\n    axs[1].imshow(detections)\n    axs[1].set_title(\"Ground truth\")\n    axs[1].axis(\"off\")\n    \n    model.eval()\n    with torch.no_grad():\n        preds = model([img.to(DEVICE)])[0]\n\n    axs[2].imshow(img.cpu().numpy().transpose((1,2,0)))\n    \n    for mask in preds['masks'].cpu().detach():\n        box = get_box(mask[[0]])\n        boxes = np.logical_or(boxes, draw_box(box))\n        \n    l = pd.Series(preds['labels'].cpu().numpy()).value_counts()\n    lstr = \"\"\n    for i in l.index:\n        lstr += f\"{l[i]}x{i} \"\n    mask_threshold = MASK_THRESHOLD_DICT[l.sort_values().index[-1]]\n    pred_masks = combine_masks(get_filtered_masks(preds), 0.5)\n        \n       \n    detections = combine_masks_boxes(pred_masks, boxes)\n    score = iou_map([masks],[pred_masks])\n    axs[2].imshow(detections)\n    axs[2].set_title(f\"Predictions | IoU score: {score:.2f}\")\n    axs[2].axis(\"off\")\n    plt.show()\nanalyze_sample(model, train_dataset, random.randint(1, 500))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"analyze_sample(model, train_dataset, random.randint(1, 500))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CellTestDataset(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    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        return {'image': image, 'image_id': image_id}","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds_test = CellTestDataset(TEST_PATH, transforms=get_transform(train=False))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nsubmission = []\nfor sample in ds_test:\n    img = sample['image']\n    image_id = sample['image_id']\n    with torch.no_grad():\n        result = model([img.to(DEVICE)])[0]\n    \n    previous_masks = []\n    for i, mask in enumerate(result[\"masks\"]):\n        score = result[\"scores\"][i].cpu().item()\n        mask = mask.cpu().numpy()\n        # Keep only highly likely pixels\n        binary_mask = mask > MASK_THRESHOLD\n        binary_mask = remove_overlapping_pixels(binary_mask, previous_masks)\n        previous_masks.append(binary_mask)\n        rle = rle_encoding(binary_mask)\n        submission.append((image_id, rle))\n    \n    # Add empty prediction if no RLE was generated for this image\n    all_images_ids = [image_id for image_id, rle in submission]\n    if image_id not in all_images_ids:\n        submission.append((image_id, \"\"))\n\ndf_sub = pd.DataFrame(submission, columns=['id', 'predicted'])\ndf_sub.to_csv(\"submission.csv\", index=False)\ndf_sub.head()","metadata":{},"outputs":[],"execution_count":null}]}