{"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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":5,"nbformat":4,"cells":[{"id":"cf6f0cda","cell_type":"markdown","source":"## 📦 Imports","metadata":{}},{"id":"ee7923b6","cell_type":"code","source":"import os, numpy as np, torch\nimport torchvision\nfrom torchvision.models.detection import maskrcnn_resnet50_fpn\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.transforms import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom sklearn.metrics import jaccard_score, f1_score\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom tqdm import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:46.475244Z","iopub.execute_input":"2025-06-20T12:37:46.475446Z","iopub.status.idle":"2025-06-20T12:37:55.812134Z","shell.execute_reply.started":"2025-06-20T12:37:46.475428Z","shell.execute_reply":"2025-06-20T12:37:55.811342Z"}},"outputs":[],"execution_count":1},{"id":"f0e7a6c0","cell_type":"markdown","source":"## ⚙️ Configuration","metadata":{}},{"id":"414bd681","cell_type":"code","source":"DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nNUM_CLASSES = 2\nEPOCHS = 10\nBATCH_SIZE = 2\nSAVE_DIR = \"./output\"\nos.makedirs(SAVE_DIR, exist_ok=True)\nIMG_DIR = \"../input/sartorius-cell-instance-segmentation/train\"\nMASK_DIR = \"./masks\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:55.813519Z","iopub.execute_input":"2025-06-20T12:37:55.813842Z","iopub.status.idle":"2025-06-20T12:37:55.896606Z","shell.execute_reply.started":"2025-06-20T12:37:55.813824Z","shell.execute_reply":"2025-06-20T12:37:55.895829Z"}},"outputs":[],"execution_count":2},{"id":"40136a46","cell_type":"markdown","source":"## 📂 Dataset Loader","metadata":{}},{"id":"7827faad","cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, image_dir, mask_dir, transforms=None):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.file_names = sorted(os.listdir(mask_dir))\n        self.transforms = transforms\n\n    def __getitem__(self, idx):\n        fname = self.file_names[idx]\n        img_path = os.path.join(self.image_dir, fname)\n        mask_path = os.path.join(self.mask_dir, fname)\n\n        img = Image.open(img_path).convert(\"RGB\")\n        mask = Image.open(mask_path).convert(\"L\")\n        img = F.to_tensor(img)\n        mask = np.array(mask)\n        obj_ids = np.unique(mask)[1:]  # 忽略背景0\n\n        masks = mask == obj_ids[:, None, None]\n        boxes = []\n        for m in masks:\n            pos = np.where(m)\n            xmin, xmax = pos[1].min(), pos[1].max()\n            ymin, ymax = pos[0].min(), pos[0].max()\n            boxes.append([xmin, ymin, xmax, ymax])\n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n        labels = torch.ones((len(obj_ids),), dtype=torch.int64)\n\n        target = {\"boxes\": boxes, \"labels\": labels, \"masks\": torch.as_tensor(masks, dtype=torch.uint8)}\n        return img, target\n\n    def __len__(self):\n        return len(self.file_names)\n\ndataset = CellDataset(IMG_DIR, MASK_DIR)\ndata_loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True, collate_fn=lambda x: tuple(zip(*x)))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:55.897487Z","iopub.execute_input":"2025-06-20T12:37:55.89783Z","iopub.status.idle":"2025-06-20T12:37:55.989839Z","shell.execute_reply.started":"2025-06-20T12:37:55.897811Z","shell.execute_reply":"2025-06-20T12:37:55.988686Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mFileNotFoundError\u001b[0m                         Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_35/1513789385.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m     33\u001b[0m         \u001b[0;32mreturn\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfile_names\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     34\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 35\u001b[0;31m \u001b[0mdataset\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mCellDataset\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mIMG_DIR\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mMASK_DIR\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     36\u001b[0m \u001b[0mdata_loader\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mDataLoader\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mbatch_size\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mBATCH_SIZE\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mshuffle\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mTrue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mcollate_fn\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mlambda\u001b[0m \u001b[0mx\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mtuple\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mzip\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_35/1513789385.py\u001b[0m in \u001b[0;36m__init__\u001b[0;34m(self, image_dir, mask_dir, transforms)\u001b[0m\n\u001b[1;32m      3\u001b[0m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mimage_dir\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mimage_dir\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      4\u001b[0m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmask_dir\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmask_dir\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 5\u001b[0;31m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfile_names\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0msorted\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mos\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlistdir\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmask_dir\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      6\u001b[0m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtransforms\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtransforms\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      7\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mFileNotFoundError\u001b[0m: [Errno 2] No such file or directory: './masks'"],"ename":"FileNotFoundError","evalue":"[Errno 2] No such file or directory: './masks'","output_type":"error"}],"execution_count":3},{"id":"96312b14","cell_type":"markdown","source":"## 🧠 Model","metadata":{}},{"id":"089c22d5","cell_type":"code","source":"model = maskrcnn_resnet50_fpn(pretrained=True)\nin_features = model.roi_heads.box_predictor.cls_score.in_features\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\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)\nmodel.to(DEVICE)\n\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.Adam(params, lr=1e-4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:55.990404Z","iopub.status.idle":"2025-06-20T12:37:55.990724Z","shell.execute_reply.started":"2025-06-20T12:37:55.990523Z","shell.execute_reply":"2025-06-20T12:37:55.990538Z"}},"outputs":[],"execution_count":null},{"id":"8fc28e01","cell_type":"markdown","source":"## 🚀 Training + Evaluation","metadata":{}},{"id":"8f5b1ae6","cell_type":"code","source":"train_loss, val_iou, val_f1 = [], [], []\n\nfor epoch in tqdm(range(EPOCHS)):\n    model.train()\n    total_loss = 0.0\n    for imgs, targets in data_loader:\n        imgs = list(img.to(DEVICE) for img in imgs)\n        targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n\n        loss_dict = model(imgs, targets)\n        loss = sum(loss for loss in loss_dict.values())\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n\n    # eval on batch\n    model.eval()\n    with torch.no_grad():\n        imgs, targets = next(iter(data_loader))\n        imgs = list(img.to(DEVICE) for img in imgs)\n        targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n        output = model(imgs)\n\n        y_true = targets[0][\"masks\"][0].cpu().numpy().astype(int).flatten()\n        y_pred = output[0][\"masks\"][0, 0].cpu().numpy()\n        y_pred = (y_pred > 0.5).astype(int).flatten()\n\n        iou = jaccard_score(y_true, y_pred)\n        f1 = f1_score(y_true, y_pred)\n\n    print(f\"Epoch {epoch+1}/{EPOCHS} | Loss={total_loss:.4f} | IoU={iou:.4f} | F1={f1:.4f}\")\n    train_loss.append(total_loss)\n    val_iou.append(iou)\n    val_f1.append(f1)\n\ntorch.save(model.state_dict(), f\"{SAVE_DIR}/model.pth\")\npd.DataFrame({\"Loss\": train_loss, \"IoU\": val_iou, \"F1\": val_f1}).to_csv(f\"{SAVE_DIR}/metrics.csv\", index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:55.991833Z","iopub.status.idle":"2025-06-20T12:37:55.99215Z","shell.execute_reply.started":"2025-06-20T12:37:55.992002Z","shell.execute_reply":"2025-06-20T12:37:55.992021Z"}},"outputs":[],"execution_count":null},{"id":"51ea5e36","cell_type":"markdown","source":"## 📊 Loss / IoU / F1 Curve","metadata":{}},{"id":"a9f91368","cell_type":"code","source":"plt.figure(figsize=(12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(train_loss); plt.title(\"Loss\")\nplt.subplot(1, 2, 2)\nplt.plot(val_iou, label=\"IoU\")\nplt.plot(val_f1, label=\"F1\")\nplt.title(\"Validation\")\nplt.legend()\nplt.savefig(f\"{SAVE_DIR}/curves.png\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T12:37:55.993274Z","iopub.status.idle":"2025-06-20T12:37:55.993596Z","shell.execute_reply.started":"2025-06-20T12:37:55.993473Z","shell.execute_reply":"2025-06-20T12:37:55.993487Z"}},"outputs":[],"execution_count":null}]}