{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"colab":{"provenance":[],"gpuType":"A100"},"accelerator":"GPU","kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"#### Load the necessary libraries","metadata":{"id":"TnYP2GglY86n"}},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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":{"id":"F_dE8Ky66ymF"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Mount the drive","metadata":{"id":"uFzqcYaSZING"}},{"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":{"id":"ucXVP56r7rtb","outputId":"bab1b28d-4328-4322-e10d-a9a5903023a0"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Get the paths and extract the data","metadata":{"id":"Lj3NcNClZNRN"}},{"cell_type":"code","source":"data_path = \"/kaggle/input/sartorius-cell-instance-segmentation\"\n# data_path = \"/content/sartorius-cell-instance-segmentation\"","metadata":{"id":"NwekSqe2ZSUn"},"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)\n","metadata":{"id":"1bgcGmhRH61V"},"outputs":[],"execution_count":null},{"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":{"id":"Zi5kzQ1_QFcv"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CSV\ndf = pd.read_csv(train_df_path)\ndf.head()","metadata":{"id":"Sbr3PpYgR2GV","outputId":"fd76cfb5-3403-40e5-ea21-32c873585087"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Historgram for class distribution","metadata":{"id":"O-jGqZBQ_WhD"}},{"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":{"id":"Ho9Vo3H0wOyH","outputId":"f14c4305-b2b6-4060-98db-e5107308b528"},"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":{"id":"V2K1ZwkU_WHC","outputId":"e0824ebc-0329-4ba8-e0ca-4b072dc5b728"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Make the train/validation split from the train folder 80/20","metadata":{"id":"j7iSBO2BVlnb"}},{"cell_type":"code","source":"unique_ids = df['id'].unique()\nprint(f\"Total unique images: {len(unique_ids)}\")","metadata":{"id":"HZ1KLJJ6Sr3-","outputId":"932d0922-3f5b-4a24-fd0c-3df6cb13418e"},"outputs":[],"execution_count":null},{"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":{"id":"C9XvmCbPSxLy"},"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()}\")\n","metadata":{"id":"wiqcJbN2TME0","outputId":"123e943a-f3ff-4730-f829-2122fccf0aa2"},"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()\n","metadata":{"id":"fQ-oMLJwTTUJ"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Visualize image and mask","metadata":{"id":"xx56GMBmYHlh"}},{"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":{"id":"2-Ev1e37YLIH","outputId":"85f3d653-3590-4d42-ccb6-2010f2c8ad87"},"outputs":[],"execution_count":null},{"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":{"id":"lfAQyG3VYVZn"},"outputs":[],"execution_count":null},{"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":{"id":"dzEmNYzaP74w"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cell_type_dict = {\n    'shsy5y': 1,\n    'astro': 2,\n    'cort': 3\n}","metadata":{"id":"Zbwhfogn5OYe"},"outputs":[],"execution_count":null},{"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":{"id":"imRUyjxG1C7h"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))\n\nresize_factor = False","metadata":{"id":"4SBzIX5n7Dh6"},"outputs":[],"execution_count":null},{"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":{"id":"DNjS2yMhQqIF"},"outputs":[],"execution_count":null},{"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":{"id":"mZF1irL1-6Na"},"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    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":{"id":"j3iy2QLS0_H8"},"outputs":[],"execution_count":null},{"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":{"id":"JYc8spoJ9t7f","outputId":"712ba21a-bd85-4022-fe1f-5dd7bd28f740"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df.columns)","metadata":{"id":"8BzuBJarOw84","outputId":"9cbe8347-be09-4894-fb34-b06bbfdda52a"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##TRAINING","metadata":{"id":"Z1jCuHTFHZSI"}},{"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":{"id":"OvGuW7NL7m9W"},"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":{"id":"YIFSO7X-GMHa"},"outputs":[],"execution_count":null},{"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":{"id":"sDASkfyB--ty","outputId":"36471000-9354-4375-b57e-4a9d1b060147","collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Two-stage optimizer setup with warm-up and grad-clip","metadata":{"id":"Fd1Mat4cu7Yz"}},{"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":{"id":"VVpVEMrG7P3L","outputId":"479b6f35-726b-4499-8222-9e7e9dd2ebe0"},"outputs":[],"execution_count":null},{"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":{"id":"zCH8Mi2AvGnl","outputId":"50cc0b69-9da1-411f-d0c8-f447826bc04a"},"outputs":[],"execution_count":null},{"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":{"id":"hzsmuMe1pRy5"},"outputs":[],"execution_count":null},{"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":{"id":"DK2hBVF67QAT"},"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":{"id":"_Ig7my_Q7QMy","outputId":"55237075-bb84-430b-b47c-d28ac7c396a1"},"outputs":[],"execution_count":null},{"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":{"id":"WGjX5JsJe2nn","outputId":"8c348854-4749-4343-9ae2-8c6c20bcbef7"},"outputs":[],"execution_count":null},{"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":{"id":"-0TeOfll4Lfg","outputId":"75b1f3fe-7ee6-4776-addc-f1243eb448a6"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"9JbIKlEp2-3S"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##EVAL MODEL","metadata":{"id":"E4FxZX76F-so"}},{"cell_type":"code","source":"#dataset for test\nclass TestCellDataset(Dataset):\n    def __init__(self, image_dir, image_size=(512, 512), transform=None):\n        self.image_dir = image_dir\n        self.image_ids = [f.replace('.png', '') for f in os.listdir(image_dir) if f.endswith('.png')]\n        self.image_size = image_size\n        self.transform = transform or T.Compose([\n            T.Resize(image_size),\n            T.ToTensor(),\n            T.Normalize(mean=[0.485, 0.456, 0.406],\n                        std=[0.229, 0.224, 0.225])\n        ])\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        img_path = os.path.join(self.image_dir, f\"{image_id}.png\")\n        image = Image.open(img_path).convert(\"RGB\")\n        image = self.transform(image)\n        return image, image_id","metadata":{"id":"68Q1VUsMHL2f"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = TestCellDataset(test_img_path)\ntest_loader = DataLoader(test_dataset, batch_size=4, shuffle=False)","metadata":{"id":"gBV2G9WqZ1YR"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"####for MaskRCNN","metadata":{"id":"p9wqLbLsLZ5E"}},{"cell_type":"code","source":"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\n#Submission generator\ndef generate_submission_csv(model, test_loader, output_csv_path='submission.csv', device='cuda', threshold=0.5):\n    model.load_state_dict(torch.load(\"checkpoints/best_model.pth\", map_location=DEVICE))\n    model.eval()\n    submissions = []\n\n    for images, image_ids in test_loader:\n        images = list(img.to(device) for img in images)\n\n        with torch.no_grad():\n            outputs = model(images)\n\n        for i, output in enumerate(outputs):\n            image_id = image_ids[i]\n            masks = output['masks'].squeeze(1).cpu().numpy() if output['masks'].ndim == 4 else []\n\n            if len(masks) == 0:\n                submissions.append({\"id\": image_id, \"predicted\": \"\"})\n                continue\n\n            for mask in masks:\n                bin_mask = (mask > threshold).astype(np.uint8)\n                if bin_mask.sum() == 0:\n                    continue\n                rle = rle_encode(bin_mask)\n                if rle:\n                    submissions.append({\"id\": image_id, \"predicted\": rle})\n\n    df = pd.DataFrame(submissions)\n    df.to_csv(output_csv_path, index=False)\n    print(f\"Submision file created : {output_csv_path}\")\n    return df","metadata":{"id":"YfGORinOHPyq"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"VikhYzPO3jQ-"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"generate_submission_csv(model, test_loader, output_csv_path='submission.csv', device = DEVICE, threshold=MASK_THRESHOLD)\n# (model, test_loader, output_csv_path='submission.csv', device='cuda', threshold=0.5):","metadata":{"id":"-IcFQow6YLTQ","outputId":"4875a5b2-7ef0-4e92-9202-1a32d70db3bb"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###For U-NET","metadata":{"id":"rab8fsovLcu6"}},{"cell_type":"code","source":"# from skimage.measure import label\n\n# def generate_submission_csv_unet(model, test_loader, output_csv_path='submission.csv', device='cuda', threshold=0.5):\n#     model.eval()\n#     submissions = []\n\n#     for images, image_ids in test_loader:\n#         images = list(img.to(device) for img in images)\n\n#         with torch.no_grad():\n#             preds = model(torch.stack(images))  # output shape: [B, 1, H, W]\n#             preds = preds.squeeze(1).cpu().numpy()\n\n#         for i in range(len(preds)):\n#             image_id = image_ids[i]\n#             bin_mask = (preds[i] > threshold).astype(np.uint8)\n\n#             labeled = label(bin_mask)\n#             if labeled.max() == 0:\n#                 submissions.append({\"id\": image_id, \"predicted\": \"\"})\n#                 continue\n\n#             for inst_id in range(1, labeled.max() + 1):\n#                 inst_mask = (labeled == inst_id).astype(np.uint8)\n#                 rle = rle_encode(inst_mask)\n#                 if rle:\n#                     submissions.append({\"id\": image_id, \"predicted\": rle})\n\n#     df = pd.DataFrame(submissions)\n#     df.to_csv(output_csv_path, index=False)\n#     print(f\"submission file (U-Net) : {output_csv_path}\")\n#     return df\n","metadata":{"id":"ECAUeF7LLe39"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###For visualization","metadata":{"id":"o9-Ou-aUM5Vx"}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score, jaccard_score\n\ndef evaluate_instance_segmentation(preds, gts, threshold=0.5):\n    assert len(preds) == len(gts)\n    ious, dices, precisions, recalls, f1s = [], [], [], [], []\n    empty_pred_count = 0\n    y_true_all, y_pred_all = [], []\n\n    for pred_mask, gt_mask in zip(preds, gts):\n        pred_bin = (pred_mask > threshold).astype(np.uint8)\n        gt_bin = (gt_mask > 0).astype(np.uint8)\n\n        if pred_bin.sum() == 0:\n            empty_pred_count += 1\n\n        y_true_flat = gt_bin.flatten()\n        y_pred_flat = pred_bin.flatten()\n\n        ious.append(jaccard_score(y_true_flat, y_pred_flat, zero_division=0))\n        dices.append(f1_score(y_true_flat, y_pred_flat, zero_division=0))\n        precisions.append(precision_score(y_true_flat, y_pred_flat, zero_division=0))\n        recalls.append(recall_score(y_true_flat, y_pred_flat, zero_division=0))\n        f1s.append(f1_score(y_true_flat, y_pred_flat, zero_division=0))\n\n        y_true_all.extend(y_true_flat)\n        y_pred_all.extend(y_pred_flat)\n\n    cm = confusion_matrix(y_true_all, y_pred_all)\n\n    return {\n        \"mean_IoU\": np.mean(ious),\n        \"mean_Dice\": np.mean(dices),\n        \"mean_Precision\": np.mean(precisions),\n        \"mean_Recall\": np.mean(recalls),\n        \"mean_F1\": np.mean(f1s),\n        \"empty_prediction_rate\": empty_pred_count / len(preds),\n        \"confusion_matrix\": cm\n    }\n\ndef plot_confusion_matrix(cm, labels=[\"background\", \"cell\"]):\n    fig, ax = plt.subplots(figsize=(4, 4))\n    ax.matshow(cm, cmap=plt.cm.Blues, alpha=0.8)\n    for i in range(cm.shape[0]):\n        for j in range(cm.shape[1]):\n            ax.text(x=j, y=i, s=cm[i, j], va='center', ha='center')\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.xticks(ticks=range(len(labels)), labels=labels)\n    plt.yticks(ticks=range(len(labels)), labels=labels)\n    plt.title(\"Confusion Matrix\")\n    plt.tight_layout()\n    plt.show()\n\ndef test_model_on_batch(model, dataloader, device='cuda', threshold=0.5, max_images=4):\n    model.eval()\n    images, targets = next(iter(dataloader))\n    images = list(img.to(device) for img in images)\n\n    with torch.no_grad():\n        outputs = model(images)\n\n    preds = []\n    gts = []\n    for output, target in zip(outputs, targets):\n        # Predicted Masks\n        pred_masks = output['masks'].squeeze(1).cpu().numpy() if output['masks'].ndim == 4 else []\n        combined_pred = np.zeros((512, 512))\n        for mask in pred_masks:\n            combined_pred = np.maximum(combined_pred, mask)\n        preds.append(combined_pred)\n\n        # Masks ground truth\n        gt_masks = target['masks'].cpu().numpy()\n        combined_gt = np.zeros((512, 512))\n        for mask in gt_masks:\n            combined_gt = np.maximum(combined_gt, mask)\n        gts.append(combined_gt)\n\n    results = evaluate_instance_segmentation(preds, gts, threshold)\n\n    for k, v in results.items():\n        if k != \"confusion_matrix\":\n            print(f\"{k}: {v:.4f}\")\n    plot_confusion_matrix(results[\"confusion_matrix\"])\n\n    return results\n","metadata":{"id":"lTzhsplAF_h1"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test_dataset = TestCellDataset(image_dir='chemin/vers/test/', image_size=(512, 512))\n# test_loader = DataLoader(test_dataset, batch_size=4, shuffle=False)\n\n\"\"\"import model\nfrom torchvision.models.detection import maskrcnn_resnet50_fpn for example\"\"\"\n#model = maskrcnn_resnet50_fpn(pretrained=True)\n#model.to('cuda')","metadata":{"id":"icmbZxiFHhdq","outputId":"ccbd8bcd-c105-4d0a-c47c-b2182ca450d9"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Submission file choose depend on your model\ngenerate_submission_csv(model, test_loader, output_csv_path='submission.csv', device='cuda')\n# generate_submission_csv_unet(model, test_loader, output_csv_path='submission.csv', device='cuda')\n\n#Visualization\n# test_model_on_batch(model, test_loader, device='cuda')","metadata":{"id":"miAaC-5HGAUv","outputId":"f8287690-10b7-406b-97db-07429f5991ee"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_instance_segmentation(preds, gts, threshold=0.5):\n    assert len(preds) == len(gts)\n    ious, dices, precisions, recalls, f1s = [], [], [], [], []\n    empty_pred_count = 0\n    y_true_all, y_pred_all = [], []\n\n    for pred_mask, gt_mask in zip(preds, gts):\n        pred_bin = (pred_mask > threshold).astype(np.uint8)\n        gt_bin = (gt_mask > 0).astype(np.uint8)\n        if pred_bin.sum() == 0:\n            empty_pred_count += 1\n        y_true_flat = gt_bin.flatten()\n        y_pred_flat = pred_bin.flatten()\n        ious.append(jaccard_score(y_true_flat, y_pred_flat, zero_division=0))\n        dices.append(f1_score(y_true_flat, y_pred_flat, zero_division=0))\n        precisions.append(precision_score(y_true_flat, y_pred_flat, zero_division=0))\n        recalls.append(recall_score(y_true_flat, y_pred_flat, zero_division=0))\n        f1s.append(f1_score(y_true_flat, y_pred_flat, zero_division=0))\n        y_true_all.extend(y_true_flat)\n        y_pred_all.extend(y_pred_flat)\n\n    cm = confusion_matrix(y_true_all, y_pred_all)\n    return {\n        \"mean_IoU\": np.mean(ious),\n        \"mean_Dice\": np.mean(dices),\n        \"mean_Precision\": np.mean(precisions),\n        \"mean_Recall\": np.mean(recalls),\n        \"mean_F1\": np.mean(f1s),\n        \"empty_prediction_rate\": empty_pred_count / len(preds),\n        \"confusion_matrix\": cm\n    }\n\n# 5. Affichage matrix\n\ndef plot_confusion_matrix(cm, labels=[\"background\", \"cell\"]):\n    fig, ax = plt.subplots(figsize=(4, 4))\n    ax.matshow(cm, cmap=plt.cm.Blues, alpha=0.8)\n    for i in range(cm.shape[0]):\n        for j in range(cm.shape[1]):\n            ax.text(x=j, y=i, s=cm[i, j], va='center', ha='center')\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.xticks(ticks=range(len(labels)), labels=labels)\n    plt.yticks(ticks=range(len(labels)), labels=labels)\n    plt.title(\"Confusion Matrix\")\n    plt.tight_layout()\n    plt.show()\n\n# 6. Test batch visuel (val_loader) pour Mask R-CNN ou UNet\n\ndef test_model_on_batch_unet(model, dataloader, device='cuda', threshold=0.5):\n    model.eval()\n    images, targets = next(iter(dataloader))\n    images = list(img.to(device) for img in images)\n\n    with torch.no_grad():\n        preds = model(torch.stack(images)).squeeze(1).cpu().numpy()\n\n    preds_bin = [(p > threshold).astype(np.uint8) for p in preds]\n    gts = [(t['masks'].sum(0).cpu().numpy() > 0).astype(np.uint8) for t in targets]\n\n    results = evaluate_instance_segmentation(preds_bin, gts, threshold)\n    for k, v in results.items():\n        if k != \"confusion_matrix\":\n            print(f\"{k}: {v:.4f}\")\n    plot_confusion_matrix(results[\"confusion_matrix\"])\n    return results","metadata":{"id":"Tb0xmvnPMfIP"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"5u8-qG1DHbxG"},"outputs":[],"execution_count":null}]}