{"metadata":{"kernelspec":{"display_name":"Python 3","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.11.11"},"accelerator":"GPU","colab":{"gpuType":"A100","provenance":[]},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":4059.360018,"end_time":"2025-05-29T13:16:47.892574","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-05-29T12:09:08.532556","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"a3f96067","cell_type":"markdown","source":"#### Load the necessary libraries","metadata":{"id":"TnYP2GglY86n","papermill":{"duration":0.007432,"end_time":"2025-05-29T12:09:12.692051","exception":false,"start_time":"2025-05-29T12:09:12.684619","status":"completed"},"tags":[]}},{"id":"f94543d9","cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:12.705719Z","iopub.status.busy":"2025-05-29T12:09:12.705478Z","iopub.status.idle":"2025-05-29T12:09:16.925159Z","shell.execute_reply":"2025-05-29T12:09:16.924519Z"},"papermill":{"duration":4.227985,"end_time":"2025-05-29T12:09:16.926616","exception":false,"start_time":"2025-05-29T12:09:12.698631","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b596a915","cell_type":"code","source":"from google.colab import drive\nimport os\nimport zipfile\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nfrom torch.utils.data import random_split\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport random\nimport torch\nimport collections\nimport torchvision\nimport time\nimport matplotlib.patches as patches\nimport cv2\nfrom torchvision.transforms import functional as F\n\n# from torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import MaskRCNN, fasterrcnn_resnet50_fpn\nfrom torchvision.models.detection.backbone_utils import resnet_fpn_backbone\n# from torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\n\nfrom functools import partial","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:16.940446Z","iopub.status.busy":"2025-05-29T12:09:16.940125Z","iopub.status.idle":"2025-05-29T12:09:24.283562Z","shell.execute_reply":"2025-05-29T12:09:24.282452Z"},"id":"F_dE8Ky66ymF","papermill":{"duration":7.351962,"end_time":"2025-05-29T12:09:24.285257","exception":false,"start_time":"2025-05-29T12:09:16.933295","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"15847f8a","cell_type":"markdown","source":"#### Mount the drive","metadata":{"id":"uFzqcYaSZING","papermill":{"duration":0.007932,"end_time":"2025-05-29T12:09:24.302421","exception":false,"start_time":"2025-05-29T12:09:24.294489","status":"completed"},"tags":[]}},{"id":"6c4e7a50","cell_type":"code","source":"#file path in google drive\n#create a shortcut in your Drive.\n# from google.colab import drive\n# drive.mount('/content/drive')\n\n# !cp -r /content/drive/MyDrive/sartorius-cell-instance-segmentation /content/","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.320928Z","iopub.status.busy":"2025-05-29T12:09:24.320166Z","iopub.status.idle":"2025-05-29T12:09:24.324244Z","shell.execute_reply":"2025-05-29T12:09:24.323395Z"},"id":"ucXVP56r7rtb","outputId":"bab1b28d-4328-4322-e10d-a9a5903023a0","papermill":{"duration":0.016689,"end_time":"2025-05-29T12:09:24.325448","exception":false,"start_time":"2025-05-29T12:09:24.308759","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"490459a3","cell_type":"markdown","source":"#### Get the paths and extract the data","metadata":{"id":"Lj3NcNClZNRN","papermill":{"duration":0.006417,"end_time":"2025-05-29T12:09:24.338166","exception":false,"start_time":"2025-05-29T12:09:24.331749","status":"completed"},"tags":[]}},{"id":"cf8b426b","cell_type":"code","source":"data_path = \"/kaggle/input/sartorius-cell-instance-segmentation\"\n# data_path = \"/content/sartorius-cell-instance-segmentation\"","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.351387Z","iopub.status.busy":"2025-05-29T12:09:24.350962Z","iopub.status.idle":"2025-05-29T12:09:24.354152Z","shell.execute_reply":"2025-05-29T12:09:24.353490Z"},"id":"NwekSqe2ZSUn","papermill":{"duration":0.010823,"end_time":"2025-05-29T12:09:24.355219","exception":false,"start_time":"2025-05-29T12:09:24.344396","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c6d49bf1","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)\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.368916Z","iopub.status.busy":"2025-05-29T12:09:24.368293Z","iopub.status.idle":"2025-05-29T12:09:24.371630Z","shell.execute_reply":"2025-05-29T12:09:24.371089Z"},"id":"1bgcGmhRH61V","papermill":{"duration":0.011109,"end_time":"2025-05-29T12:09:24.372671","exception":false,"start_time":"2025-05-29T12:09:24.361562","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c321fa42","cell_type":"code","source":"# get the img and mask paths\n# TRAIN_CSV = f\"{data_dir}/train.csv\"\n# TRAIN_PATH = f\"{data_dir}/train\"\n# TEST_PATH = f\"{data_dir}/test\"\n\ntest_img_path = f\"{data_path}/test\"\ntrain_img_path = f\"{data_path}/train\"\ntrain_df_path = f\"{data_path}/train.csv\" # annotations (image IDs + RLE masks)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.385587Z","iopub.status.busy":"2025-05-29T12:09:24.385378Z","iopub.status.idle":"2025-05-29T12:09:24.388753Z","shell.execute_reply":"2025-05-29T12:09:24.388080Z"},"id":"Zi5kzQ1_QFcv","papermill":{"duration":0.010916,"end_time":"2025-05-29T12:09:24.389772","exception":false,"start_time":"2025-05-29T12:09:24.378856","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1fe5591f","cell_type":"code","source":"# CSV\ndf = pd.read_csv(train_df_path)\ndf.head()","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.403054Z","iopub.status.busy":"2025-05-29T12:09:24.402860Z","iopub.status.idle":"2025-05-29T12:09:24.924404Z","shell.execute_reply":"2025-05-29T12:09:24.923666Z"},"id":"Sbr3PpYgR2GV","outputId":"fd76cfb5-3403-40e5-ea21-32c873585087","papermill":{"duration":0.52967,"end_time":"2025-05-29T12:09:24.925795","exception":false,"start_time":"2025-05-29T12:09:24.396125","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"84ae1be7","cell_type":"markdown","source":"#### Historgram for class distribution","metadata":{"id":"O-jGqZBQ_WhD","papermill":{"duration":0.007371,"end_time":"2025-05-29T12:09:24.942995","exception":false,"start_time":"2025-05-29T12:09:24.935624","status":"completed"},"tags":[]}},{"id":"f1b3939b","cell_type":"code","source":"df = pd.read_csv(train_df_path)\ndf.head()\ncell_type_counts = df['cell_type'].value_counts()\nprint(cell_type_counts)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:24.960871Z","iopub.status.busy":"2025-05-29T12:09:24.960628Z","iopub.status.idle":"2025-05-29T12:09:25.320684Z","shell.execute_reply":"2025-05-29T12:09:25.319636Z"},"id":"Ho9Vo3H0wOyH","outputId":"f14c4305-b2b6-4060-98db-e5107308b528","papermill":{"duration":0.371423,"end_time":"2025-05-29T12:09:25.322238","exception":false,"start_time":"2025-05-29T12:09:24.950815","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"f149e9e1","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":{"execution":{"iopub.execute_input":"2025-05-29T12:09:25.344824Z","iopub.status.busy":"2025-05-29T12:09:25.344590Z","iopub.status.idle":"2025-05-29T12:09:25.624058Z","shell.execute_reply":"2025-05-29T12:09:25.623345Z"},"id":"V2K1ZwkU_WHC","outputId":"e0824ebc-0329-4ba8-e0ca-4b072dc5b728","papermill":{"duration":0.291361,"end_time":"2025-05-29T12:09:25.625530","exception":false,"start_time":"2025-05-29T12:09:25.334169","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"45975076","cell_type":"markdown","source":"#### Make the train/validation split from the train folder 80/20","metadata":{"id":"j7iSBO2BVlnb","papermill":{"duration":0.006921,"end_time":"2025-05-29T12:09:25.645490","exception":false,"start_time":"2025-05-29T12:09:25.638569","status":"completed"},"tags":[]}},{"id":"8f24dbb8","cell_type":"code","source":"unique_ids = df['id'].unique()\nprint(f\"Total unique images: {len(unique_ids)}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:25.660760Z","iopub.status.busy":"2025-05-29T12:09:25.660124Z","iopub.status.idle":"2025-05-29T12:09:25.671195Z","shell.execute_reply":"2025-05-29T12:09:25.670449Z"},"id":"HZ1KLJJ6Sr3-","outputId":"932d0922-3f5b-4a24-fd0c-3df6cb13418e","papermill":{"duration":0.019962,"end_time":"2025-05-29T12:09:25.672485","exception":false,"start_time":"2025-05-29T12:09:25.652523","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"f8c0f2c4","cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# 80% train, 20% validation\ntrain_ids, val_ids = train_test_split(\n    unique_ids,\n    test_size=0.2,\n    random_state=42,\n    shuffle=True\n)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:25.693183Z","iopub.status.busy":"2025-05-29T12:09:25.692973Z","iopub.status.idle":"2025-05-29T12:09:26.874179Z","shell.execute_reply":"2025-05-29T12:09:26.873633Z"},"id":"C9XvmCbPSxLy","papermill":{"duration":1.190278,"end_time":"2025-05-29T12:09:26.875473","exception":false,"start_time":"2025-05-29T12:09:25.685195","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"fa615c54","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()}\")\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:26.891405Z","iopub.status.busy":"2025-05-29T12:09:26.890692Z","iopub.status.idle":"2025-05-29T12:09:26.921484Z","shell.execute_reply":"2025-05-29T12:09:26.920573Z"},"id":"wiqcJbN2TME0","outputId":"123e943a-f3ff-4730-f829-2122fccf0aa2","papermill":{"duration":0.039752,"end_time":"2025-05-29T12:09:26.922598","exception":false,"start_time":"2025-05-29T12:09:26.882846","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bfb2f00c","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()\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:26.979654Z","iopub.status.busy":"2025-05-29T12:09:26.979182Z","iopub.status.idle":"2025-05-29T12:09:27.741585Z","shell.execute_reply":"2025-05-29T12:09:27.740986Z"},"id":"fQ-oMLJwTTUJ","papermill":{"duration":0.771782,"end_time":"2025-05-29T12:09:27.742826","exception":false,"start_time":"2025-05-29T12:09:26.971044","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"20515cd8","cell_type":"markdown","source":"#### Visualize image and mask","metadata":{"id":"xx56GMBmYHlh","papermill":{"duration":0.007462,"end_time":"2025-05-29T12:09:27.757867","exception":false,"start_time":"2025-05-29T12:09:27.750405","status":"completed"},"tags":[]}},{"id":"4266ed9f","cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport os\nfrom PIL import Image\n\nWIDTH = 704\nHEIGHT = 520\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], train_df, train_img_path)\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:27.773112Z","iopub.status.busy":"2025-05-29T12:09:27.772918Z","iopub.status.idle":"2025-05-29T12:09:28.433319Z","shell.execute_reply":"2025-05-29T12:09:28.432649Z"},"id":"2-Ev1e37YLIH","outputId":"85f3d653-3590-4d42-ccb6-2010f2c8ad87","papermill":{"duration":0.677013,"end_time":"2025-05-29T12:09:28.442015","exception":false,"start_time":"2025-05-29T12:09:27.765002","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d8b261bf","cell_type":"code","source":"#Transform\nimport torchvision.transforms as T\n\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\nNORMALIZE = False\n# transform = T.Compose([\n#     T.Resize((512, 512)),\n#     T.ToTensor(),\n#     T.Normalize(\n#       mean=RESNET_MEAN,\n#       std=RESNET_STD)\n# ])\n\n# for not pretrained model\n# import albumentations as A\n# from albumentations.pytorch import ToTensorV2\n\n# transform = A.Compose(\n#   [\n#     A.RandomCrop(512,512),\n#     A.HorizontalFlip(p=0.5),\n#     A.Rotate(limit=30, p=0.5),\n#     A.Normalize(mean=RESNET_MEAN, std=RESNET_STD),\n#     ToTensorV2()\n#   ],\n#   bbox_params=A.BboxParams(\n#     format='pascal_voc',     # for[xmin,ymin,xmax,ymax]\n#     label_fields=['labels']\n#   )\n# )","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.473438Z","iopub.status.busy":"2025-05-29T12:09:28.473140Z","iopub.status.idle":"2025-05-29T12:09:28.477403Z","shell.execute_reply":"2025-05-29T12:09:28.476655Z"},"id":"lfAQyG3VYVZn","papermill":{"duration":0.021235,"end_time":"2025-05-29T12:09:28.478653","exception":false,"start_time":"2025-05-29T12:09:28.457418","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e691fad9","cell_type":"code","source":"class 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\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    # Data augmentation for train\n    if train:\n        transforms.append(HorizontalFlip(0.5))\n        transforms.append(VerticalFlip(0.5))\n\n    return Compose(transforms)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.508383Z","iopub.status.busy":"2025-05-29T12:09:28.508103Z","iopub.status.idle":"2025-05-29T12:09:28.515974Z","shell.execute_reply":"2025-05-29T12:09:28.515513Z"},"id":"dzEmNYzaP74w","papermill":{"duration":0.024313,"end_time":"2025-05-29T12:09:28.516947","exception":false,"start_time":"2025-05-29T12:09:28.492634","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ef47cf67","cell_type":"code","source":"cell_type_dict = {\n    'shsy5y': 1,\n    'astro': 2,\n    'cort': 3\n}","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.549848Z","iopub.status.busy":"2025-05-29T12:09:28.549465Z","iopub.status.idle":"2025-05-29T12:09:28.553224Z","shell.execute_reply":"2025-05-29T12:09:28.552458Z"},"id":"Zbwhfogn5OYe","papermill":{"duration":0.02302,"end_time":"2025-05-29T12:09:28.554483","exception":false,"start_time":"2025-05-29T12:09:28.531463","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"cb77a31f","cell_type":"code","source":"def rle_decode(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\nclass CellDataset(Dataset):\n    def __init__(self, image_dir, df, transforms=None, resize=False):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.df = df\n\n        self.should_resize = resize is not False\n        if self.should_resize:\n            self.height = int(HEIGHT * resize)\n            self.width = int(WIDTH * resize)\n            print(\"image size used:\", self.height, self.width)\n        else:\n            self.height = HEIGHT\n            self.width = WIDTH\n\n        self.image_info = collections.defaultdict(dict)\n        temp_df = self.df.groupby([\"id\", \"cell_type\"])['annotation'].agg(lambda x: list(x)).reset_index()\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': list(row[\"annotation\"]),\n                    'cell_type': cell_type_dict[row[\"cell_type\"]]\n                    }\n\n    def get_box(self, a_mask):\n        ''' Get the bounding box of a given 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        return [xmin, ymin, xmax, ymax]\n\n    def __getitem__(self, idx):\n        ''' Get the image and the target'''\n\n        img_path = self.image_info[idx][\"image_path\"]\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n\n        if self.should_resize:\n            img = cv2.resize(img, (self.width, self.height))\n\n        info = self.image_info[idx]\n\n        n_objects = len(info['annotations'])\n        masks = np.zeros((len(info['annotations']), self.height, self.width), dtype=np.uint8)\n        boxes = []\n        labels = []\n        for i, annotation in enumerate(info['annotations']):\n            a_mask = rle_decode(annotation, (HEIGHT, WIDTH))\n\n            if self.should_resize:\n                a_mask = cv2.resize(a_mask, (self.width, self.height))\n\n            a_mask = np.array(a_mask) > 0\n            masks[i, :, :] = a_mask\n\n            boxes.append(self.get_box(a_mask))\n\n        # labels\n        labels = [int(info[\"cell_type\"]) for _ in range(n_objects)]\n        #labels = [1 for _ in range(n_objects)]\n\n\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\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        # This is the required target for the Mask R-CNN\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\n        return img, target\n\n    def __len__(self):\n        return len(self.image_info)\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.587922Z","iopub.status.busy":"2025-05-29T12:09:28.587440Z","iopub.status.idle":"2025-05-29T12:09:28.601681Z","shell.execute_reply":"2025-05-29T12:09:28.601102Z"},"id":"imRUyjxG1C7h","papermill":{"duration":0.032205,"end_time":"2025-05-29T12:09:28.603002","exception":false,"start_time":"2025-05-29T12:09:28.570797","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"caec894d","cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))\n\nresize_factor = False","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.636539Z","iopub.status.busy":"2025-05-29T12:09:28.635824Z","iopub.status.idle":"2025-05-29T12:09:28.639501Z","shell.execute_reply":"2025-05-29T12:09:28.638926Z"},"id":"4SBzIX5n7Dh6","papermill":{"duration":0.021624,"end_time":"2025-05-29T12:09:28.640739","exception":false,"start_time":"2025-05-29T12:09:28.619115","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"aca751c2","cell_type":"code","source":"train_dataset = CellDataset(train_img_path, train_df, resize=resize_factor, transforms=get_transform(train=True))\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True, pin_memory=True,\n                      num_workers=2, collate_fn=collate_fn)\n\nval_dataset = CellDataset(train_img_path, val_df, resize=resize_factor, transforms=get_transform(train=False))\nval_loader = DataLoader(val_dataset, batch_size=2, shuffle=True, pin_memory=True,\n                    num_workers=2, collate_fn=collate_fn)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.672799Z","iopub.status.busy":"2025-05-29T12:09:28.672517Z","iopub.status.idle":"2025-05-29T12:09:28.742066Z","shell.execute_reply":"2025-05-29T12:09:28.741297Z"},"id":"DNjS2yMhQqIF","papermill":{"duration":0.086203,"end_time":"2025-05-29T12:09:28.743284","exception":false,"start_time":"2025-05-29T12:09:28.657081","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"90b6be5d","cell_type":"code","source":"# train_dataset = CellDataset(\n#     df=train_df,\n#     image_dir=train_img_path,\n#     transform=transform,\n#     augment=True,\n#     cls_map=cls_map\n# )\n# # DataLoader for training\n# train_loader = DataLoader(\n#     train_dataset,\n#     batch_size=4,\n#     shuffle=True,\n#     num_workers=2,  # increase if possible\n#     collate_fn=collate_fn\n# )\n\n# # DataLoader for validation\n# val_dataset = CellDataset(\n#     df=val_df,\n#     image_dir=train_img_path,\n#     transform=transform,\n#     augment=False,\n#     cls_map=cls_map\n# )\n# val_loader = DataLoader(\n#     val_dataset,\n#     batch_size=4,\n#     shuffle=False,\n#     collate_fn=collate_fn\n# )","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.772213Z","iopub.status.busy":"2025-05-29T12:09:28.771814Z","iopub.status.idle":"2025-05-29T12:09:28.775071Z","shell.execute_reply":"2025-05-29T12:09:28.774523Z"},"id":"mZF1irL1-6Na","papermill":{"duration":0.018642,"end_time":"2025-05-29T12:09:28.776110","exception":false,"start_time":"2025-05-29T12:09:28.757468","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"29d99377","cell_type":"code","source":"def show_batch(images, targets, max_images=4):\n    plt.figure(figsize=(16, 4 * max_images))\n\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n\n    for i in range(min(max_images, len(images))):\n        image = images[i].permute(1, 2, 0).cpu().numpy()\n        image = std * image + mean  # full channel-wise denorm\n        image = np.clip(image, 0, 1)\n\n        masks = targets[i][\"masks\"].cpu().numpy()\n        combined_mask = np.zeros(image.shape[:2], dtype=np.uint8)\n\n        for mask in masks:\n            if mask.shape != image.shape[:2]:\n                mask = cv2.resize(mask, (image.shape[1], image.shape[0]), interpolation=cv2.INTER_NEAREST)\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":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.804632Z","iopub.status.busy":"2025-05-29T12:09:28.804240Z","iopub.status.idle":"2025-05-29T12:09:28.811284Z","shell.execute_reply":"2025-05-29T12:09:28.810753Z"},"id":"j3iy2QLS0_H8","papermill":{"duration":0.022223,"end_time":"2025-05-29T12:09:28.812292","exception":false,"start_time":"2025-05-29T12:09:28.790069","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d3ef7148","cell_type":"code","source":"# Récupère un batch et affiche\nimages, targets = next(iter(train_loader))\n\n# image = image.resize(self.image_size, resample=Image.BILINEAR)  # image is 512x512\nshow_batch(images, targets)\n\n# Image shape: (520, 704, 3)\n# Mask shape: (520, 704)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:28.856589Z","iopub.status.busy":"2025-05-29T12:09:28.856011Z","iopub.status.idle":"2025-05-29T12:09:31.878224Z","shell.execute_reply":"2025-05-29T12:09:31.877350Z"},"id":"JYc8spoJ9t7f","outputId":"712ba21a-bd85-4022-fe1f-5dd7bd28f740","papermill":{"duration":3.060088,"end_time":"2025-05-29T12:09:31.891411","exception":false,"start_time":"2025-05-29T12:09:28.831323","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"da4f1f09","cell_type":"code","source":"print(df.columns)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:31.947585Z","iopub.status.busy":"2025-05-29T12:09:31.947258Z","iopub.status.idle":"2025-05-29T12:09:31.951953Z","shell.execute_reply":"2025-05-29T12:09:31.951359Z"},"id":"8BzuBJarOw84","outputId":"9cbe8347-be09-4894-fb34-b06bbfdda52a","papermill":{"duration":0.033399,"end_time":"2025-05-29T12:09:31.952930","exception":false,"start_time":"2025-05-29T12:09:31.919531","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"74b56f3b","cell_type":"markdown","source":"##TRAINING","metadata":{"id":"Z1jCuHTFHZSI","papermill":{"duration":0.024804,"end_time":"2025-05-29T12:09:32.002984","exception":false,"start_time":"2025-05-29T12:09:31.978180","status":"completed"},"tags":[]}},{"id":"8e52c375","cell_type":"code","source":"import torch.nn as nn\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nMOMENTUM = 0.9\n# early stop patience 10/ epochs ~\n# LR = 0.005\n# LR = 1e-3 ~17 epoch/7 early stop\nWEIGHT_DECAY = 0.0005\nMASK_THRESHOLD = 0.5 #0.05\nPATIENCE = 3\nNUM_CLASSES = 3\nWIDTH = 704\nHEIGHT = 520\nUSE_SCHEDULER = True\n\nBATCH_SIZE = 2\n# EPOCHS = 20\nEPOCHS = 50\n# USE_SCHEDULER = False\nBOX_DETECTIONS_PER_IMG = 539\n\nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:32.053635Z","iopub.status.busy":"2025-05-29T12:09:32.053400Z","iopub.status.idle":"2025-05-29T12:09:32.058016Z","shell.execute_reply":"2025-05-29T12:09:32.057465Z"},"id":"OvGuW7NL7m9W","papermill":{"duration":0.031182,"end_time":"2025-05-29T12:09:32.059064","exception":false,"start_time":"2025-05-29T12:09:32.027882","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b8e73848","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":{"execution":{"iopub.execute_input":"2025-05-29T12:09:32.110321Z","iopub.status.busy":"2025-05-29T12:09:32.109828Z","iopub.status.idle":"2025-05-29T12:09:32.114746Z","shell.execute_reply":"2025-05-29T12:09:32.114213Z"},"id":"YIFSO7X-GMHa","papermill":{"duration":0.031661,"end_time":"2025-05-29T12:09:32.115785","exception":false,"start_time":"2025-05-29T12:09:32.084124","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"87529aff","cell_type":"code","source":"# model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=True,\n#                                                                    box_detections_per_img=BOX_DETECTIONS_PER_IMG,\n#                                                                    image_mean=RESNET_MEAN,\n#                                                                    image_std=RESNET_STD)\n\n# in_features = model.roi_heads.box_predictor.cls_score.in_features\n# model.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES+1)\n# in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n# hidden_layer = 256\n# model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES+1)\n\n# GroupNorm instead of BatchNorm for stability at small batch\ngn = lambda num_channels: nn.GroupNorm(num_groups=32, num_channels=num_channels, eps=1e-5)\n\nbackbone = resnet_fpn_backbone(\n    'resnet50',\n    pretrained=False,\n    norm_layer=gn\n)\n\n\n# model = MaskRCNN(backbone=backbone, num_classes=NUM_CLASSES+1)\n\nmodel = MaskRCNN(\n    backbone=backbone,\n    num_classes=NUM_CLASSES+1,\n    image_mean=list(RESNET_MEAN),\n    image_std=list(RESNET_STD)\n)\n\nmodel.to(DEVICE)\n\n# for param in model.parameters():\n#     param.requires_grad = True\n\n# model.train()","metadata":{"collapsed":true,"execution":{"iopub.execute_input":"2025-05-29T12:09:32.167165Z","iopub.status.busy":"2025-05-29T12:09:32.166962Z","iopub.status.idle":"2025-05-29T12:09:32.807202Z","shell.execute_reply":"2025-05-29T12:09:32.806586Z"},"id":"sDASkfyB--ty","jupyter":{"outputs_hidden":true},"outputId":"36471000-9354-4375-b57e-4a9d1b060147","papermill":{"duration":0.667121,"end_time":"2025-05-29T12:09:32.808390","exception":false,"start_time":"2025-05-29T12:09:32.141269","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a2e28b6f","cell_type":"markdown","source":"# Two-stage optimizer setup with warm-up and grad-clip","metadata":{"id":"Fd1Mat4cu7Yz","papermill":{"duration":0.025548,"end_time":"2025-05-29T12:09:32.860360","exception":false,"start_time":"2025-05-29T12:09:32.834812","status":"completed"},"tags":[]}},{"id":"5f1cc8ba","cell_type":"code","source":"# scale down the default Kaiming init in RPN, box & mask heads\ndef init_head(m):\n    if isinstance(m, torch.nn.Conv2d):\n        torch.nn.init.kaiming_normal_(m.weight, a=1)\n        m.weight.data *= 0.1\n        if m.bias is not None:\n            m.bias.data.zero_()\n\nmodel.rpn.head.apply(init_head)\nmodel.roi_heads.box_head.apply(init_head)\nmodel.roi_heads.mask_head.apply(init_head)\n\n# freeze all BatchNorm layers\nfor m in model.backbone.modules():\n    if isinstance(m, torch.nn.BatchNorm2d):\n        m.eval()\n        for p in m.parameters():\n            p.requires_grad = False\n\n# at first train heads only for 5 epochs\nfor name, p in model.backbone.named_parameters():\n    p.requires_grad = False\n\nhead_params = [p for p in model.parameters() if p.requires_grad]\nopt_stage1 = torch.optim.SGD(\n    head_params, lr=1e-4, momentum=0.9, weight_decay=WEIGHT_DECAY\n)\n\n# Warm-up scheduler: linearly ramp from 0→1 over the first 500 steps\ndef warmup_lambda(step):\n    return min((step + 1) / 500, 1.0)\n\nwarmup_sched = torch.optim.lr_scheduler.LambdaLR(opt_stage1, warmup_lambda)\n\n# params = [p for p in model.parameters() if p.requires_grad]\n# optimizer = torch.optim.SGD(\n#     params,\n#     lr=LR,\n#     momentum=MOMENTUM,\n#     weight_decay=WEIGHT_DECAY\n# )\n\n# # lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n# lr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,\n#                                                           mode='min',\n#                                                           factor=0.5,\n#                                                           patience=PATIENCE,\n#                                                           verbose=True)\nn_batches, n_batches_val = len(train_loader), len(val_loader)\n\nvalidation_mask_losses = []\ntrain_losses = []\nval_losses = []\n\nfor epoch in range(1, 10+1):\n  print(f\"Starting epoch {epoch} of 10\")\n  time_start = time.time()\n  epoch_loss = 0.0\n  loss_mask_accum = 0.0\n  loss_classifier_accum = 0.0\n  for images, targets in train_loader:\n    images  = [img.to(DEVICE) for img 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_dict.values())\n\n    opt_stage1.zero_grad()\n    loss.backward()\n    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n    opt_stage1.step()\n    warmup_sched.step()\n\n    loss_mask = loss_dict['loss_mask'].item()\n    epoch_loss += loss.item()\n    loss_mask_accum += loss_mask\n    loss_classifier_accum += loss_dict['loss_classifier'].item()\n\n  train_loss = epoch_loss / n_batches\n  train_loss_mask = loss_mask_accum / n_batches\n  train_loss_classifier = loss_classifier_accum / n_batches\n\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(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  print(f\"[Epoch {epoch} / 10] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / 10] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / 10] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / 10] Train loss: {train_loss:7.3f}. Val loss: {val_loss:7.3f}\")\n  print(f\"Time for epoch {epoch}: {epoch_time:.2f} seconds\")","metadata":{"execution":{"iopub.execute_input":"2025-05-29T12:09:32.912452Z","iopub.status.busy":"2025-05-29T12:09:32.912158Z","iopub.status.idle":"2025-05-29T12:30:37.810104Z","shell.execute_reply":"2025-05-29T12:30:37.809137Z"},"id":"VVpVEMrG7P3L","outputId":"479b6f35-726b-4499-8222-9e7e9dd2ebe0","papermill":{"duration":1264.956939,"end_time":"2025-05-29T12:30:37.842815","exception":false,"start_time":"2025-05-29T12:09:32.885876","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"83e623df","cell_type":"code","source":"for p in model.backbone.parameters():\n    p.requires_grad = True\n\nbackbone_params, head_params = [], []\nfor name, p in model.named_parameters():\n    (backbone_params if \"backbone\" in name else head_params).append(p)\n\nopt_stage2 = torch.optim.SGD([\n    {\"params\": head_params, \"lr\": 1e-3},\n    {\"params\": backbone_params, \"lr\": 1e-5},\n], momentum=0.9, weight_decay=WEIGHT_DECAY)\n\n# Step LR or ReduceLROnPlateau on val loss?\nlr_scheduler = torch.optim.lr_scheduler.StepLR(opt_stage2, step_size=10, gamma=0.1)\n\n\nvalidation_mask_losses = []\ntrain_losses = []\nval_losses = []\n\nbest_val_loss = float('inf')\nepochs_no_improve = 0\nearly_stop_patience = 10\nearly_stop_min_delta = 1e-4\n\nos.makedirs(\"checkpoints\", exist_ok=True)\n\nfor epoch in range(11, EPOCHS+1):\n  print(f\"Starting epoch {epoch} of {EPOCHS}\")\n  time_start = time.time()\n  epoch_loss = 0.0\n  loss_mask_accum = 0.0\n  loss_classifier_accum = 0.0\n  for images, targets in train_loader:\n    images  = [img.to(DEVICE) for img 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_dict.values())\n\n    opt_stage2.zero_grad()\n    loss.backward()\n    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n    opt_stage2.step()\n    warmup_sched.step()\n\n    loss_mask = loss_dict['loss_mask'].item()\n    epoch_loss += loss.item()\n    loss_mask_accum += loss_mask\n    loss_classifier_accum += loss_dict['loss_classifier'].item()\n\n  if USE_SCHEDULER:\n    lr_scheduler.step()\n\n  train_loss = epoch_loss / n_batches\n  train_loss_mask = loss_mask_accum / n_batches\n  train_loss_classifier = loss_classifier_accum / n_batches\n\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(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  # Early stopping\n  if val_loss < best_val_loss - early_stop_min_delta:\n    best_val_loss = val_loss\n    epochs_no_improve = 0\n    best_model_path = \"checkpoints/best_model.pth\"\n    torch.save(model.state_dict(), best_model_path)\n    print(f\"New best model saved: {best_model_path} with val loss {val_loss:.4f}\")\n  else:\n    epochs_no_improve += 1\n    print(f\"No improvement in val loss for {epochs_no_improve} epochs\")\n\n  if epochs_no_improve >= early_stop_patience:\n    print(f\"Early stopping at epoch {epoch}, no improvemenet in {early_stop_patience} epochs.\")\n    break\n\n  print(f\"[Epoch {epoch} / {EPOCHS}] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / {EPOCHS}] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / {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":{"execution":{"iopub.execute_input":"2025-05-29T12:30:37.896966Z","iopub.status.busy":"2025-05-29T12:30:37.896697Z","iopub.status.idle":"2025-05-29T13:16:36.486787Z","shell.execute_reply":"2025-05-29T13:16:36.485823Z"},"id":"zCH8Mi2AvGnl","outputId":"50cc0b69-9da1-411f-d0c8-f447826bc04a","papermill":{"duration":2758.64801,"end_time":"2025-05-29T13:16:36.517379","exception":false,"start_time":"2025-05-29T12:30:37.869369","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"9a88bbbd","cell_type":"code","source":"# for i, (img, tgt) in enumerate(zip(images, targets)):\n#     print(f\"\\n--- Sample {i} ---\")\n#     print(\"Image shape:\", img.shape)\n#     print(\"Image stats:\",\n#           f\"min={img.min().item():.3f}\",\n#           f\"max={img.max().item():.3f}\",\n#           f\"mean={img.mean().item():.3f}\",\n#           f\"std={img.std().item():.3f}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-29T13:16:36.583124Z","iopub.status.busy":"2025-05-29T13:16:36.582839Z","iopub.status.idle":"2025-05-29T13:16:36.586937Z","shell.execute_reply":"2025-05-29T13:16:36.586281Z"},"id":"hzsmuMe1pRy5","papermill":{"duration":0.038536,"end_time":"2025-05-29T13:16:36.588166","exception":false,"start_time":"2025-05-29T13:16:36.549630","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e02eae04","cell_type":"code","source":"# os.makedirs(\"checkpoints\", exist_ok=True)\n\n# validation_mask_losses = []\n# train_losses = []\n# val_losses = []\n\n# for epoch in range(1, EPOCHS + 1):\n#     print(f\"Starting epoch {epoch} of {EPOCHS}\")\n\n#     time_start = time.time()\n#     epoch_loss = 0.0\n#     loss_mask_accum = 0.0\n#     loss_classifier_accum = 0.0\n#     for batch_idx, (images, targets) in enumerate(train_loader, 1):\n\n#         images = list(image.to(DEVICE) for image in images)\n#         # images = [img.to(DEVICE, memory_format=torch.channels_last) for img 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#         epoch_loss += loss.item()\n#         loss_mask_accum += loss_mask\n#         loss_classifier_accum += loss_dict['loss_classifier'].item()\n\n#         if batch_idx % 500 == 0:\n#             print(f\"[Batch {batch_idx:3d} / {n_batches:3d}] Batch train loss: {loss.item():7.3f}. Mask-only loss: {loss_mask:7.3f}.\")\n\n#     if USE_SCHEDULER:\n#         lr_scheduler.step()\n\n#     train_loss = epoch_loss / n_batches\n#     train_loss_mask = loss_mask_accum / n_batches\n#     train_loss_classifier = loss_classifier_accum / n_batches\n\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(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#     checkpoint_path = f\"checkpoints/maskrcnn_epoch_{epoch}.pth\"\n#     torch.save(model.state_dict(), checkpoint_path)\n\n#     print(f\"[Epoch {epoch} / {EPOCHS}] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n#     print(f\"[Epoch {epoch} / {EPOCHS}] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n#     print(f\"[Epoch {epoch} / {EPOCHS}] Train loss: {train_loss:7.3f}. Val loss: {val_loss:7.3f}\")\n#     print(f\"Time for epoch {epoch}: {epoch_time:.2f} seconds\")\n#     print(f\"Saved checkpoint: {checkpoint_path}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-29T13:16:36.654577Z","iopub.status.busy":"2025-05-29T13:16:36.654290Z","iopub.status.idle":"2025-05-29T13:16:36.658909Z","shell.execute_reply":"2025-05-29T13:16:36.658292Z"},"id":"DK2hBVF67QAT","papermill":{"duration":0.039828,"end_time":"2025-05-29T13:16:36.660109","exception":false,"start_time":"2025-05-29T13:16:36.620281","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b4a9fe87","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":{"execution":{"iopub.execute_input":"2025-05-29T13:16:36.720452Z","iopub.status.busy":"2025-05-29T13:16:36.720200Z","iopub.status.idle":"2025-05-29T13:16:36.899334Z","shell.execute_reply":"2025-05-29T13:16:36.898648Z"},"id":"_Ig7my_Q7QMy","outputId":"55237075-bb84-430b-b47c-d28ac7c396a1","papermill":{"duration":0.208639,"end_time":"2025-05-29T13:16:36.900500","exception":false,"start_time":"2025-05-29T13:16:36.691861","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b1798c0c","cell_type":"code","source":"import matplotlib.patches as patches\n\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport numpy as np\nimport torch\n\ndef show_batch_with_preds(images, targets, outputs, max_images=4, mask_threshold=0.5):\n    plt.figure(figsize=(16, 4 * max_images))\n\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n\n    for i in range(min(max_images, len(images))):\n        image = images[i].cpu()\n        target = targets[i]\n        output = outputs[i]\n\n        img_np = image.permute(1, 2, 0).numpy()\n        img_np = np.clip((img_np * std) + mean, 0, 1)\n\n        # Ground Truth\n        gt_masks = target[\"masks\"].cpu().numpy()\n        gt_combined_mask = np.zeros(img_np.shape[:2], dtype=np.uint8)\n        for mask in gt_masks:\n            gt_combined_mask = np.maximum(gt_combined_mask, mask.squeeze().astype(np.uint8))\n\n        ax1 = plt.subplot(max_images, 2, 2 * i + 1)\n        ax1.imshow(img_np)\n        ax1.imshow(gt_combined_mask, alpha=0.4, cmap='Blues')\n        ax1.set_title(\"Ground Truth\")\n        ax1.axis('off')\n\n        for box in target[\"boxes\"].cpu():\n            x1, y1, x2, y2 = box.tolist()\n            rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n                                     linewidth=2, edgecolor='blue', facecolor='none')\n            ax1.add_patch(rect)\n\n        # Predictions\n        scores = output[\"scores\"].cpu()\n        keep = scores > mask_threshold\n\n        pred_masks = output[\"masks\"][keep].cpu().numpy()\n        pred_boxes = output[\"boxes\"][keep].cpu().numpy()\n\n        pred_combined_mask = np.zeros(img_np.shape[:2], dtype=np.uint8)\n        for mask in pred_masks:\n            pred_combined_mask = np.maximum(pred_combined_mask, (mask.squeeze() > 0.5).astype(np.uint8))\n\n        ax2 = plt.subplot(max_images, 2, 2 * i + 2)\n        ax2.imshow(img_np)\n        ax2.imshow(pred_combined_mask, alpha=0.4, cmap='Reds')\n        ax2.set_title(\"Prediction\")\n        ax2.axis('off')\n\n        for box in pred_boxes:\n            x1, y1, x2, y2 = box\n            rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n                                     linewidth=2, edgecolor='red', facecolor='none')\n            ax2.add_patch(rect)\n    print([output['labels'] for output in outputs])\n    print(outputs[0]['scores'])\n    plt.tight_layout()\n    plt.show()\n\nmodel.eval()\nwith torch.no_grad():\n    outputs = model([img.to(DEVICE) for img in images])  # make sure images is a list\nshow_batch_with_preds(images, targets, outputs, max_images=3)\n","metadata":{"execution":{"iopub.execute_input":"2025-05-29T13:16:36.961140Z","iopub.status.busy":"2025-05-29T13:16:36.960904Z","iopub.status.idle":"2025-05-29T13:16:38.590065Z","shell.execute_reply":"2025-05-29T13:16:38.589417Z"},"id":"WGjX5JsJe2nn","outputId":"8c348854-4749-4343-9ae2-8c6c20bcbef7","papermill":{"duration":1.664593,"end_time":"2025-05-29T13:16:38.596542","exception":false,"start_time":"2025-05-29T13:16:36.931949","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"66a941c2","cell_type":"code","source":"show_batch_with_preds(\n    images, targets, outputs,\n    max_images=3,\n    mask_threshold=0.2  # or even 0.1 if needed\n)","metadata":{"execution":{"iopub.execute_input":"2025-05-29T13:16:38.687838Z","iopub.status.busy":"2025-05-29T13:16:38.687552Z","iopub.status.idle":"2025-05-29T13:16:40.120415Z","shell.execute_reply":"2025-05-29T13:16:40.119772Z"},"id":"-0TeOfll4Lfg","outputId":"75b1f3fe-7ee6-4776-addc-f1243eb448a6","papermill":{"duration":1.48505,"end_time":"2025-05-29T13:16:40.126719","exception":false,"start_time":"2025-05-29T13:16:38.641669","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c9b13309","cell_type":"markdown","source":"##EVAL MODEL","metadata":{"id":"E4FxZX76F-so","papermill":{"duration":0.053223,"end_time":"2025-05-29T13:16:40.345147","exception":false,"start_time":"2025-05-29T13:16:40.291924","status":"completed"},"tags":[]}},{"id":"5d9b3afa","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":{},"outputs":[],"execution_count":null},{"id":"ea69671b","cell_type":"code","source":"class 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":{},"outputs":[],"execution_count":null},{"id":"ee0a3139","cell_type":"code","source":"model.load_state_dict(torch.load(\"checkpoints/best_model.pth\", map_location=DEVICE))\nmodel.to(DEVICE).eval()\n\ntest_dataset = CellTestDataset(test_img_path, transforms=get_transform(train=False))\ntest_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{},"outputs":[],"execution_count":null},{"id":"35809cf9","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":{},"outputs":[],"execution_count":null},{"id":"65ae29fe","cell_type":"code","source":"print(f\"{len(test_dataset)}images found in test dataset\")\n\nsubmission = []\nfor batch in test_loader:\n    images   = batch[\"image\"].to(DEVICE)\n    img_ids  = batch[\"image_id\"]\n\n    with torch.no_grad():\n        outputs = model(images)\n\n    for out, image_id in zip(outputs, img_ids):\n        H, W = out[\"masks\"].shape[-2:]\n        occupied = np.zeros((H, W), dtype=bool)\n\n        for mask_tensor, score in zip(out[\"masks\"], out[\"scores\"]):\n            if score.item() < 0.3:\n                continue\n\n            mask_np = mask_tensor[0].cpu().numpy()\n            bin_mask = mask_np > 0.5\n\n            # remove overlaps\n            clean_mask = remove_overlapping_pixels(bin_mask, occupied)\n            if clean_mask.sum() == 0:\n                continue\n\n            # mark those pixels as occupied\n            occupied |= clean_mask\n\n            rle = rle_encoding(clean_mask.astype(np.uint8))\n            submission.append((image_id, rle))\n\ndf = pd.DataFrame(submission, columns=[\"id\", \"predicted\"])\ndf.to_csv(\"submission.csv\", index=False)\nprint(f\"{len(df)} mask‐rows in submission file\")","metadata":{},"outputs":[],"execution_count":null}]}