{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3"},"language_info":{"name":"python"},"colab":{"provenance":[],"gpuType":"T4"},"accelerator":"GPU","kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12012286,"sourceType":"datasetVersion","datasetId":7557074},{"sourceId":12013409,"sourceType":"datasetVersion","datasetId":7557819}],"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":"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 torch.nn as nn\nfrom torchvision.ops import misc as misc_nn_ops\nimport cv2\nfrom torchvision.transforms import functional as F\nimport torchvision.transforms as T\nfrom torchvision.models.detection import MaskRCNN, fasterrcnn_resnet50_fpn\nfrom torchvision.models.detection.backbone_utils import resnet_fpn_backbone\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nimport torch.nn as nn\nfrom functools import partial\nfrom sklearn.model_selection import train_test_split","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":"e2741682-833c-4c6f-e3bc-8211b9068c1e"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Get the paths and extract the data","metadata":{"id":"Lj3NcNClZNRN"}},{"cell_type":"code","source":"# data_path = \"/content/sartorius-cell-instance-segmentation\"\ndata_path = \"/kaggle/input/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\ntest_img_path = f\"{data_path}/test\"\ntrain_img_path = f\"{data_path}/train\"\ntrain_df_path = f\"{data_path}/train.csv\"","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":"9b332d62-b86e-4ee4-a14a-183155503fea"},"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":"cb412091-f979-4c75-aefb-9b909c64abcb"},"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":"05c2a0dc-3acc-4aa3-de01-923f10db4061"},"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":"641dec32-dab5-4186-ad3b-8ff62fed6d37"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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":"e267fd1d-a5f2-4bb3-ca84-2567bf4d4608"},"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":"def 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    print(f\"Max: {np.max(combined_mask)}, min: {np.min(combined_mask)}\")\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":"6b052b6d-4948-40ad-accc-d8dabf0565a4"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Transform\n\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\nNORMALIZE = False\nresize_factor = False\ncell_type_dict = {\"astro\": 1, \"cort\": 2, \"shsy5y\": 3}\n# mask_threshold_dict = {1: 0.55, 2: 0.75, 3:  0.6}\n# min_score_dict = {1: 0.55, 2: 0.75, 3: 0.5}\n# resize_factor = False\nWIDTH = 704\nHEIGHT = 520\nBATCH_SIZE = 2\n\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":"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))","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=BATCH_SIZE, 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=BATCH_SIZE, shuffle=False, pin_memory=True,\n                    num_workers=2, collate_fn=collate_fn)\n#shuffle false for val","metadata":{"id":"DNjS2yMhQqIF"},"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\nshow_batch(images, targets)","metadata":{"id":"JYc8spoJ9t7f","outputId":"3bc4ff02-d90d-48be-e7a6-40fe7d8b2fe2"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df.columns)","metadata":{"id":"8BzuBJarOw84","outputId":"114757dc-da48-47ce-a143-3d8a7c33bfe4"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##TRAINING","metadata":{"id":"Z1jCuHTFHZSI"}},{"cell_type":"code","source":"MOMENTUM = 0.9\nWEIGHT_DECAY = 0.0005\nWEIGHT_DECAY2 = 1e-4\nMASK_THRESHOLD = 0.5\nPATIENCE = 3\nNUM_CLASSES = 3\nWIDTH = 704\nHEIGHT = 520\nUSE_SCHEDULER = False\n# USE_SCHEDULER = True\n# BATCH_SIZE = 4 (CUDA out of memory)\nEPOCHS = 50\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":"# for kaggle!!!\nfrom torchvision.models import resnet50\nfrom torchvision.models.detection.backbone_utils import _resnet_fpn_extractor\nfrom torchvision.models.detection import MaskRCNN\n\nresnet = resnet50(norm_layer=torch.nn.BatchNorm2d, weights=None)\n\n#weights\nstate_dict = torch.load(\"/kaggle/input/resnet50-weights/resnet50-0676ba61.pth\")\nresnet.load_state_dict(state_dict)  # Full load, no strict=False\n\n#fpn\nbackbone = _resnet_fpn_extractor(resnet, trainable_layers=3)\nbackbone.out_channels = 256  # important!\n\nmodel = MaskRCNN(\n    backbone=backbone,\n    num_classes=NUM_CLASSES + 1  # your foreground + background\n)\n\n#trained model before\ntrained_state = torch.load(\"/kaggle/input/best-model-pth/best_model.pth\", map_location=DEVICE)\nmodel.load_state_dict(trained_state)  # strict=True should now work\nmodel.to(DEVICE)","metadata":{"id":"sDASkfyB--ty","outputId":"56861caf-1b04-479f-a13a-775583f0803c","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":"\"\"\"\nscale down the default Kaiming init in RPN, box & mask heads\n-stabilize early training weights = too large at the beginning of training\n\"\"\"\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# train heads 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-3, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY\n)\n\n# Warm-up schedule: step 500 full LR\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\nEPOCH_ = 5\n\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, EPOCH_+1):\n  print(f\"Starting epoch {epoch} of {EPOCH_}\")\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          #output = model(images)\n          #print(output)\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} / {EPOCH_}] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / {EPOCH_}] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n  print(f\"[Epoch {epoch} / {EPOCH_}] 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":"c9c38341-4332-4088-fcda-317209ae1423"},"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":"CfA7vuqu0vRP"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# best_model_path = \"best_model.pth\"\n# torch.save(model.state_dict(), best_model_path)","metadata":{"id":"ka42dApbGptX"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model.load_state_dict(torch.load(\"best_model_pre.pth\", map_location=DEVICE))\n\nfor 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\n# opt_stage2 = torch.optim.SGD([\n#     {\"params\": head_params, \"lr\": 1e-1},\n#     {\"params\": backbone_params, \"lr\": 1e-4},\n# ], momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\n\nopt_stage2 = torch.optim.AdamW([\n    {\"params\": head_params, \"lr\": 1e-4},\n    {\"params\": backbone_params, \"lr\": 1e-5},\n], weight_decay=WEIGHT_DECAY2)\n\n# Step LR or ReduceLROnPlateau on val loss for SGD StepLR\n# lr_scheduler = torch.optim.lr_scheduler.StepLR(opt_stage2, step_size=10, gamma=0.1)\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    opt_stage2, mode='min', factor=0.5, patience=3, threshold=1e-4, verbose=True\n)\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\nEPOCHS = 20\n\nfor epoch in range(EPOCH_+1, 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(val_losses)\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 USE_SCHEDULER:\n    lr_scheduler.step(val_loss)\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"},"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"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.patches as patches\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport numpy as np\nimport torch\n\nmodel.load_state_dict(torch.load(\"checkpoints/best_model.pth\", map_location=DEVICE))\n\ndef remove_overlapping_pixels(mask, other_masks):\n\n    for other_mask in other_masks:\n\n        if np.sum(np.logical_and(mask, other_mask)) > 0:\n            mask[np.logical_and(mask, other_mask)] = 0\n\n    return mask\n\n#mask_threshold_dict = {1: 0.05, 2: 0.05, 3:  0.06}\nmask_threshold_dict = {1: 0.55, 2: 0.75, 3:  0.6}\nmin_score_dict = {1: 0.55, 2: 0.75, 3: 0.5}\n\ndef get_filtered_masks_v2(masks, boxes, scores, labels):\n    \"\"\"\n    filter masks using MIN_SCORE for mask and MAX_THRESHOLD for pixels\n    \"\"\"\n    use_masks = []\n    use_boxes = []\n\n    for i, mask in enumerate(masks):\n        #print(\"check mask # \", i)\n        label = labels[i]\n        #binary_mask = mask > mask_threshold_dict[label]\n        binary_mask = remove_overlapping_pixels(mask, use_masks)\n\n        #fig, ax = plt.subplots(nrows=1, ncols=1, figsize=(10,10))\n        #ax.imshow(binary_mask)\n        #print(boxes[i])\n        #x1, y1, x2, y2 = boxes[i]\n        #rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n        #                             linewidth=2, edgecolor='red', facecolor='none')\n        #ax.add_patch(rect)\n        #plt.show()\n\n        if np.any(binary_mask) and np.sum(np.array(binary_mask, dtype=int)) > 100:\n            use_masks.append(binary_mask)\n            use_boxes.append(boxes[i])\n            #print(f\"not all pixels eliminated! {np.sum(np.array(binary_mask, dtype=int))}\")\n        #else:\n            #print(\"all pixels eliminated\")\n\n    return use_masks, use_boxes\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        mask_threshold = 0.3\n        keep = scores > mask_threshold\n\n        pred_masks = output[\"masks\"][keep].cpu().numpy()\n        pred_boxes = output[\"boxes\"][keep].cpu().numpy()\n        pred_labels = output[\"labels\"][keep].cpu().numpy()\n        pred_scores = output[\"scores\"][keep].cpu().numpy()\n        print(pred_masks.shape)\n        binary_masks = []\n\n        for id , mask in enumerate(pred_masks):\n            label = pred_labels[id]\n            #print(label)\n            mask = mask[0]\n            binary_mask = mask > mask_threshold_dict[label]\n\n            binary_masks.append(binary_mask)\n\n            #fig, ax = plt.subplots(nrows=1, ncols=1, figsize=(10,10))\n            #ax.imshow(binary_mask)\n            #print(pred_boxes[id])\n            #x1, y1, x2, y2 = pred_boxes[id]\n            #rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n            #                         linewidth=2, edgecolor='red', facecolor='none')\n            #ax.add_patch(rect)\n            #plt.show()\n\n        print(\"before check \", len(binary_masks))\n\n\n        binary_masks, pred_boxes = get_filtered_masks_v2(binary_masks, pred_boxes, pred_scores, pred_labels)\n        print(\"after check \", len(pred_masks))\n        pred_combined_mask = np.zeros(img_np.shape[:2], dtype=np.uint8)\n\n        pred_combined_ADD_mask = np.zeros(img_np.shape[:2], dtype=np.uint8)\n\n        for id, mask in enumerate(binary_masks):\n            #pred_combined_mask += mask\n            mask = np.array(mask, dtype=int)\n            #x1, y1, x2, y2 = pred_boxes[id]\n            #print(mask[int(y1):int(y2), int(x1):int(x2)])\n            pred_combined_mask = np.maximum(pred_combined_mask, (mask.squeeze() > 0.5).astype(np.uint8))\n            pred_combined_ADD_mask += (mask.squeeze() > 0.5).astype(np.uint8)\n\n            fig, ax = plt.subplots(nrows=1, ncols=1, figsize=(10,10))\n            ax.imshow(mask)\n            print(pred_boxes[id])\n            x1, y1, x2, y2 = pred_boxes[id]\n            rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n                                     linewidth=2, edgecolor='red', facecolor='none')\n\n            ax.add_patch(rect)\n            plt.show()\n            print(np.sum(mask))\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\n        #ax3 = plt.subplot(max_images, 2, 2 * i + 2)\n        #ax3.imshow(img_np)\n        #ax3.imshow(pred_combined_ADD_mask, alpha=0.4, cmap='Blues')\n        #ax3.set_title(\"Check for overlaps between masks\")\n        #ax1.axis('off')\n\n    print([output['labels'] for output in outputs])\n    print(outputs[0]['scores'])\n    print(outputs[0]['masks'])\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"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef compute_iou(labels, y_pred, verbose=0):\n    \"\"\"\n    Computes the IoU for instance labels and predictions.\n\n    Args:\n        labels (np array): Labels.\n        y_pred (np array): predictions\n\n    Returns:\n        np array: IoU matrix, of size true_objects x pred_objects.\n    \"\"\"\n\n    true_objects = len(np.unique(labels))\n    pred_objects = len(np.unique(y_pred))\n\n    if verbose:\n        print(\"Number of true objects: {}\".format(true_objects))\n        print(\"Number of predicted objects: {}\".format(pred_objects))\n\n    # Compute intersection between all objects\n    intersection = np.histogram2d(\n        labels.flatten(), y_pred.flatten(), bins=(true_objects, pred_objects)\n    )[0]\n\n    # Compute areas (needed for finding the union between all objects)\n    area_true = np.histogram(labels, bins=true_objects)[0]\n    area_pred = np.histogram(y_pred, bins=pred_objects)[0]\n    area_true = np.expand_dims(area_true, -1)\n    area_pred = np.expand_dims(area_pred, 0)\n\n    # Compute union\n    union = area_true + area_pred - intersection\n    intersection = intersection[1:, 1:] # exclude background\n    union = union[1:, 1:]\n    union[union == 0] = 1e-9\n    iou = intersection / union\n\n    return iou\n\ndef precision_at(threshold, iou):\n    \"\"\"\n    Computes the precision at a given threshold.\n\n    Args:\n        threshold (float): Threshold.\n        iou (np array): IoU matrix.\n\n    Returns:\n        int: Number of true positives,\n        int: Number of false positives,\n        int: Number of false negatives.\n    \"\"\"\n    matches = iou > threshold\n    true_positives = np.sum(matches, axis=1) == 1  # Correct objects\n    false_positives = np.sum(matches, axis=0) == 0  # Missed objects\n    false_negatives = np.sum(matches, axis=1) == 0  # Extra objects\n    tp, fp, fn = (\n        np.sum(true_positives),\n        np.sum(false_positives),\n        np.sum(false_negatives),\n    )\n    return tp, fp, fn\n\ndef iou_map(truths, preds, verbose=0):\n    \"\"\"\n    Computes the metric for the competition.\n    Masks contain the segmented pixels where each object has one value associated,\n    and 0 is the background.\n\n    Args:\n        truths (list of masks): Ground truths.\n        preds (list of masks): Predictions.\n        verbose (int, optional): Whether to print infos. Defaults to 0.\n\n    Returns:\n        float: mAP.\n    \"\"\"\n    ious = [compute_iou(truth, pred, verbose) for truth, pred in zip(truths, preds)]\n\n    if verbose:\n        print(\"Thresh\\tTP\\tFP\\tFN\\tPrec.\")\n\n    prec = []\n    for t in np.arange(0.5, 1.0, 0.05):\n        tps, fps, fns = 0, 0, 0\n        for iou in ious:\n            tp, fp, fn = precision_at(t, iou)\n            tps += tp\n            fps += fp\n            fns += fn\n\n        p = tps / (tps + fps + fns)\n        prec.append(p)\n\n        if verbose:\n            print(\"{:1.3f}\\t{}\\t{}\\t{}\\t{:1.3f}\".format(t, tps, fps, fns, p))\n\n    if verbose:\n        print(\"AP\\t-\\t-\\t-\\t{:1.3f}\".format(np.mean(prec)))\n\n    return np.mean(prec)\n","metadata":{"id":"SLeMoZVKXWS7"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cell_type_dict = {\"astro\": 1, \"cort\": 2, \"shsy5y\": 3}\n# mask_threshold_dict = {1: 0.55, 2: 0.75, 3:  0.6}\nmask_threshold_dict = {1: 0.05, 2: 0.05, 3:  0.06}\nmin_score_dict = {1: 0.55, 2: 0.75, 3: 0.5}\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\n\ndef combine_masks(masks, mask_threshold):\n    \"\"\"\n    combine masks into one image\n    \"\"\"\n    maskimg = np.zeros((HEIGHT, WIDTH))\n    # print(len(masks.shape), masks.shape)\n    for m, mask in enumerate(masks,1):\n        maskimg[mask>mask_threshold] = m\n\n    return maskimg\n\ndef combine_masks_v2(masks, mask_threshold):\n    \"\"\"\n    combine masks into one image\n    \"\"\"\n    maskimg = np.zeros((HEIGHT, WIDTH))\n    print(masks.shape)\n    # print(len(masks.shape), masks.shape)\n    for m, mask in enumerate(masks,1):\n        #print(mask[0].shape)\n        maskimg[mask[0]>mask_threshold] = m\n\n    return maskimg\n\n\ndef get_filtered_masks(pred):\n    \"\"\"\n    filter masks using MIN_SCORE for mask and MAX_THRESHOLD for pixels\n    \"\"\"\n    use_masks = []\n\n    for i, mask in enumerate(pred[\"masks\"]):\n\n        # Filter-out low-scoring results. Not tried yet.\n        scr = pred[\"scores\"][i].cpu().item()\n        label = pred[\"labels\"][i].cpu().item()\n        if scr > min_score_dict[label]:\n            mask = mask.cpu().numpy().squeeze()\n            # Keep only highly likely pixels\n            binary_mask = mask >  mask_threshold_dict[label]\n            #binary_mask = remove_overlapping_pixels(binary_mask, use_masks)\n            use_masks.append(binary_mask)\n\n    return use_masks\n\ndef analyze_train_sample(model, train_dataset, sample_index):\n\n    img, targets = train_dataset[sample_index]\n    #print(img.shape)\n    l = np.unique(targets[\"labels\"])\n    ig, ax = plt.subplots(nrows=1, ncols=3, figsize=(20,60), facecolor=\"#fefefe\")\n    ax[0].imshow(img.numpy().transpose((1,2,0)))\n    ax[0].set_title(f\"cell type {l}\")\n    ax[0].axis(\"off\")\n\n    masks = combine_masks(targets['masks'], 0.5)\n    #plt.imshow(img.numpy().transpose((1,2,0)))\n    ax[1].imshow(masks)\n    ax[1].set_title(f\"Ground truth, {len(targets['masks'])} cells\")\n    ax[1].axis(\"off\")\n\n    model.eval()\n    with torch.no_grad():\n        preds = model([img.to(DEVICE)])[0]\n    print(targets['masks'])\n    print(preds['masks'])\n\n    l = pd.Series(preds['labels'].cpu().numpy()).value_counts()\n    lstr = \"\"\n    for i in l.index:\n        lstr += f\"{l[i]}x{i} \"\n    #print(l, l.sort_values().index[-1])\n    #plt.imshow(img.cpu().numpy().transpose((1,2,0)))\n    mask_threshold = mask_threshold_dict[l.sort_values().index[-1]]\n\n    print(mask_threshold)\n    pred_masks = combine_masks(get_filtered_masks(preds), mask_threshold)\n    mask_threshold = 0.5\n    pred_masks = combine_masks_v2(preds['masks'].cpu().numpy(), mask_threshold)\n    ax[2].imshow(pred_masks)\n    ax[2].set_title(f\"Predictions, labels: {lstr}\")\n    ax[2].axis(\"off\")\n    plt.show()\n\n    #print(masks.shape, pred_masks.shape)\n    score = iou_map([masks],[pred_masks])\n    print(\"Score:\", score)\n\n\n# NOTE: It puts the model in eval mode!! Revert for re-training\nanalyze_train_sample(model, train_dataset, 20)","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)\nprint(len(test_dataset))\ntest_loader = DataLoader(test_dataset, batch_size=4, shuffle=False)\nprint(len(test_loader))","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            #print(outputs)\n\n        for i, output in enumerate(outputs):\n            image_id = image_ids[i]\n\n            scores = output[\"scores\"].cpu()\n            print(scores)\n            mask_threshold = 0.5 #0.2\n            keep = scores > mask_threshold\n            pred_scores = output[\"scores\"][keep].cpu().numpy()\n            #print(keep)\n            #masks = output[\"masks\"].squeeze(1).cpu().numpy() if output['masks'].ndim == 4 else []\n            masks = output[\"masks\"][keep].cpu().numpy()\n            boxes = output[\"boxes\"][keep].cpu().numpy()\n            labels = output[\"labels\"][keep].cpu().numpy()\n\n\n            print(masks.shape)\n            print(boxes.shape)\n            print(labels.shape)\n\n            binary_masks = []\n\n            for id , mask in enumerate(masks):\n                label = labels[id]\n                #print(label)\n                mask = mask[0]\n                binary_mask = mask > mask_threshold_dict[label]\n\n                binary_masks.append(binary_mask)\n\n\n            print(\"before check \", len(binary_masks))\n\n\n            binary_masks, pred_boxes = get_filtered_masks_v2(binary_masks, boxes, pred_scores, labels)\n            print(\"after check \", len(binary_masks))\n\n            if len(binary_masks) == 0:\n                submissions.append({\"id\": image_id, \"predicted\": \"\"})\n                continue\n\n            for mask in binary_masks:\n                bin_mask = mask.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":"generate_submission_csv(model, test_loader, output_csv_path='submission.csv', device = DEVICE, threshold=0.5)\n# (model, test_loader, output_csv_path='submission.csv', device='cuda', threshold=0.5):","metadata":{"id":"-IcFQow6YLTQ"},"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\n# 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\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\n# def 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# def 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"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Submission file choose depend on your model\n# generate_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"},"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":"# import csv\n\n# with open('sample_submission.csv', 'r') as file:\n#        csvreader = csv.reader(file)\n#        for row in csvreader:\n#            print(row)","metadata":{"id":"5u8-qG1DHbxG"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"id":"vea_uQjGVt-p"},"outputs":[],"execution_count":null}]}