{"metadata":{"kernelspec":{"display_name":"Python 3","name":"python3"},"language_info":{"name":"python"},"accelerator":"GPU","colab":{"gpuType":"A100","provenance":[]},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"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 torchvision\nimport time\nimport matplotlib.patches as patches\n\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor","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":"3bbe2ce9-46ce-4038-8c45-b4821ebfb253"},"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\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":"e504b2c1-dae0-431a-ddcd-2d7aa8458605"},"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":"eed9b66c-f8e1-47cd-95b1-3b152b212f11"},"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":"e64e9ce2-60cd-46a1-e04a-d8bf78406c35"},"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":"87e395a8-774c-48a1-de2e-20967632be92"},"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":"a03663fa-cb74-4053-c5d9-ea970d54eb05"},"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\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":"b8de6cbe-29b7-4693-b683-7bdfeeb9c1bd"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Transform\nimport torchvision.transforms as T\ntransform = T.Compose([\n    T.Resize((512, 512)),\n    T.ToTensor(),\n    T.Normalize(\n      mean=[0.485, 0.456, 0.406],\n      std=[0.229, 0.224, 0.225])\n])","metadata":{"id":"lfAQyG3VYVZn"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_decode_for_train(mask_rle, shape=(520, 704)):\n    if pd.isnull(mask_rle):\n        return np.zeros(shape, dtype=np.uint8)\n    s = mask_rle.strip().split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0::2], s[1::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)\n\nclass CellDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None, image_size=(512, 512), augment=False):\n        self.df = df\n        self.image_ids = df['id'].unique()\n        self.image_dir = image_dir\n        self.transform = transform\n        self.augment = augment\n        self.image_size = image_size\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\n        # Resize\n        if self.transform:\n            image = self.transform(image)\n\n\n        records = self.df[self.df['id'] == image_id]\n        masks = []\n        boxes = []\n\n        for _, row in records.iterrows():\n            mask = mask = rle_decode_for_train(row['annotation'])\n            mask = Image.fromarray(mask).resize(self.image_size, resample=Image.NEAREST)\n            mask = np.array(mask)\n            if mask.max() == 0:\n                continue\n            masks.append(mask)\n\n            pos = np.where(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            boxes.append([xmin, ymin, xmax, ymax])\n\n\n        if len(masks) == 0:\n            masks = torch.zeros((0, *self.image_size), dtype=torch.uint8)\n            boxes = torch.zeros((0, 4), dtype=torch.float32)\n            labels = torch.zeros((0,), dtype=torch.int64)\n        else:\n            masks = torch.tensor(np.stack(masks), dtype=torch.uint8)\n            boxes = torch.tensor(boxes, dtype=torch.float32)\n            labels = torch.ones((len(masks),), dtype=torch.int64)\n\n        target = {\n            \"boxes\": boxes,\n            \"labels\": labels,\n            \"masks\": masks,\n            \"image_id\": torch.tensor([idx])\n        }\n\n        # Simple Data Augmentation (random horizontal flip)\n        if self.augment and random.random() > 0.5:\n            image = torch.flip(image, dims=[2])  # horizontal flip\n            masks = torch.flip(masks, dims=[2])\n            boxes[:, [0, 2]] = self.image_size[1] - boxes[:, [2, 0]]  # update x coords\n            target[\"boxes\"] = boxes\n            target[\"masks\"] = masks\n\n        return image, target","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(\n    df=train_df,\n    image_dir=train_img_path,\n    transform=transform,\n    augment=True\n)\n# DataLoader for training\ntrain_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\nval_dataset = CellDataset(\n    df=val_df,\n    image_dir=train_img_path,\n    transform=transform,\n    augment=False\n)\nval_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    for i in range(min(max_images, len(images))):\n        image = images[i].permute(1, 2, 0).cpu().numpy()\n        image = np.clip(image * 0.229 + 0.485, 0, 1)  # denomarlization ImageNet\n\n        masks = targets[i][\"masks\"].cpu().numpy()\n        combined_mask = np.zeros(image.shape[:2], dtype=np.uint8)\n        for mask in masks:\n            combined_mask = np.maximum(combined_mask, mask)\n\n        plt.subplot(max_images, 2, 2 * i + 1)\n        plt.imshow(image)\n        plt.title(\"Image\")\n        plt.axis('off')\n\n        plt.subplot(max_images, 2, 2 * i + 2)\n        plt.imshow(image)\n        plt.imshow(combined_mask, alpha=0.5, cmap='viridis')\n        plt.title(\"Image + Masks\")\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"id":"j3iy2QLS0_H8"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Récupère un batch et affiche\nimages, targets = next(iter(train_loader))\nshow_batch(images, targets)","metadata":{"id":"JYc8spoJ9t7f","outputId":"a4f4353f-7577-49bd-fcee-7c6fd133291d"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##TRAINING","metadata":{"id":"Z1jCuHTFHZSI"}},{"cell_type":"code","source":"RESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nMOMENTUM = 0.9\n# LR = 0.001 #trained\nLR = 0.005\nWEIGHT_DECAY = 0.0005\nMASK_THRESHOLD = 0.5 #0.05\nPATIENCE = 3\nNUM_CLASSES = 3\nWIDTH = 704\nHEIGHT = 520\nUSE_SCHEDULER = False\n\nBATCH_SIZE = 2\nEPOCHS = 20\nUSE_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":"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\nin_features = model.roi_heads.box_predictor.cls_score.in_features\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES+1)\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\nhidden_layer = 256\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES+1)\n\n\nmodel.to(DEVICE)\n\nfor param in model.parameters():\n    param.requires_grad = True\n\nmodel.train()","metadata":{"id":"sDASkfyB--ty","outputId":"132f3c92-a6ea-4e60-b078-ed6a5518d3c7"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=LR, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\n# lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\nlr_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)","metadata":{"id":"VVpVEMrG7P3L","outputId":"1ca4a76c-920e-43de-8a44-ae336db9636a"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(\"checkpoints\", exist_ok=True)\n\nvalidation_mask_losses = []\ntrain_losses = []\nval_losses = []\n\nfor 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:2d} / {EPOCHS:2d}] Train-mask loss: {train_loss_mask:7.3f}, classifier loss {train_loss_classifier:7.3f}\")\n    print(f\"[Epoch {epoch:2d} / {EPOCHS:2d}] Val-mask loss  : {val_loss_mask:7.3f}, classifier loss {val_loss_classifier:7.3f}\")\n    print(f\"[Epoch {epoch:2d} / {EPOCHS:2d}] 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","outputId":"bfe58cde-52a2-46b8-83e5-f5e77a9f85ad"},"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":"78867e9c-50b6-4f6f-c330-5b8b9c86f706"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.patches as patches\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    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        # de-normalize for ImageNet\n        img_np = np.clip(img_np * 0.229 + 0.485, 0, 1)\n\n        # GT masks\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())\n\n        # GT plot\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        # GT boxes\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        # Prediction masks\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        pred_labels = output[\"labels\"][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)\n\n        # Plot Prediction\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        # Add predicted boxes\n        for box, label in zip(pred_boxes, pred_labels):\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            # ax2.text(x1, y1, f'Class {label}', color='white', fontsize=10,\n            #          bbox=dict(facecolor='red', alpha=0.5))\n\n    plt.tight_layout()\n    plt.show()\n\n\nmodel.eval()\nwith torch.no_grad():\n    outputs = model(images)\nshow_batch_with_preds(images, targets, outputs, max_images=3)\n","metadata":{"id":"WGjX5JsJe2nn","outputId":"1c1519e0-9b1f-464e-dd49-be61b7515a3f"},"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/maskrcnn_epoch_20.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":"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":"d8b95e00-bea0-415c-ae7e-b3bde671a8bf"},"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\n# from 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","outputId":"8b7f7ff0-6118-4174-cd1b-13f18197f121"},"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\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# # 6. Test batch visuel (val_loader) pour Mask R-CNN ou UNet\n\n# def 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}]}