{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.9","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":23823,"databundleVersionId":1920183,"sourceType":"competition"},{"sourceId":1911681,"sourceType":"datasetVersion","datasetId":1128406},{"sourceId":1934626,"sourceType":"datasetVersion","datasetId":1128710},{"sourceId":3075714,"sourceType":"datasetVersion","datasetId":849808},{"sourceId":12127290,"sourceType":"datasetVersion","datasetId":7636434},{"sourceId":12148879,"sourceType":"datasetVersion","datasetId":7651552},{"sourceId":244667974,"sourceType":"kernelVersion"},{"sourceId":433812,"sourceType":"modelInstanceVersion","modelInstanceId":353724,"modelId":375030},{"sourceId":438557,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":357800,"modelId":379137},{"sourceId":441190,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":358793,"modelId":380105}],"dockerImageVersionId":30056,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is simple mmdetection infrence script as a base line.\nTraining part can be foud [here](https://www.kaggle.com/its7171/mmdetection-for-segmentation-training).","metadata":{"papermill":{"duration":0.009854,"end_time":"2021-02-02T02:49:13.549001","exception":false,"start_time":"2021-02-02T02:49:13.539147","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install \"../input/landmark-additional-packages/timm-0.3.4-py3-none-any.whl\" # already in the system\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T07:57:33.405822Z","iopub.execute_input":"2025-06-19T07:57:33.406069Z","iopub.status.idle":"2025-06-19T07:58:12.057373Z","shell.execute_reply.started":"2025-06-19T07:57:33.406012Z","shell.execute_reply":"2025-06-19T07:58:12.056646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom torchvision import transforms\nimport timm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T07:58:12.059726Z","iopub.execute_input":"2025-06-19T07:58:12.059964Z","iopub.status.idle":"2025-06-19T07:58:13.742288Z","shell.execute_reply.started":"2025-06-19T07:58:12.059942Z","shell.execute_reply":"2025-06-19T07:58:13.741703Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 从掩膜提取 crop 图像","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nfrom PIL import Image\nfrom tqdm import tqdm\n\nmask_root = '/kaggle/input/mmdetec-mask-20-40/masks_20_40_thr05/masks_thr05'  # 替换成你的测试掩膜路径\nimage_root = '/kaggle/input/hpa-single-cell-image-classification/train'\nsave_dir = '/kaggle/working/cell_crops_test_0.2'\nos.makedirs(save_dir, exist_ok=True)\n\ndef load_rgb_image(image_id, image_root):\n    colors = ['red', 'green', 'blue']\n    imgs = []\n    for c in colors:\n        path = os.path.join(image_root, f\"{image_id}_{c}.png\")\n        if not os.path.exists(path):\n            print(f\"[跳过] 通道图像缺失: {path}\")\n            return None\n        img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        imgs.append(img)\n    return np.stack(imgs, axis=-1)\n\n# 遍历每个 image_id 文件夹\n# ✅ 只选取前 0.3% image_id 文件夹\nall_image_ids = sorted(os.listdir(mask_root))\nnum_limit = int(len(all_image_ids) * 0.2)\nselected_image_ids = all_image_ids[:num_limit]\n\nfor image_id in tqdm(selected_image_ids, desc=\"🚀 正在处理前 20% 图像\"):\n    mask_dir = os.path.join(mask_root, image_id)\n    if not os.path.isdir(mask_dir):\n        continue\n\n    rgb_img = load_rgb_image(image_id, image_root)\n\n    # 加载失败跳过\n    if rgb_img is None or len(rgb_img.shape) != 3:\n        print(f\"[跳过] 图像加载失败或维度异常：{image_id}\")\n        continue\n\n    for fname in os.listdir(mask_dir):\n        if not fname.endswith('.png'):\n            continue\n\n        try:\n            class_id = int(fname.split('_')[0].replace('class', ''))\n            cell_id = int(fname.split('_')[1].replace('cell', '').replace('.png', ''))\n        except:\n            print(f\"[跳过] 无法解析类名或细胞编号：{fname}\")\n            continue\n\n        mask_path = os.path.join(mask_dir, fname)\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n\n        if mask is None:\n            print(f\"[跳过] 掩膜文件加载失败: {mask_path}\")\n            continue\n\n        ys, xs = np.where(mask > 0)\n        if len(xs) == 0 or len(ys) == 0:\n            continue\n\n        x_min, x_max = xs.min(), xs.max()\n        y_min, y_max = ys.min(), ys.max()\n\n        crop = rgb_img[y_min:y_max+1, x_min:x_max+1]\n\n        if crop.size == 0 or len(crop.shape) != 3:\n            print(f\"[跳过] 空 crop 或维度异常：{fname}\")\n            continue\n\n        try:\n            crop_pil = Image.fromarray(cv2.cvtColor(crop, cv2.COLOR_BGR2RGB))\n        except:\n            print(f\"[跳过] PIL 转换失败：{fname}\")\n            continue\n\n        out_name = f\"{image_id}_class{class_id}_cell{cell_id}.png\"\n        crop_pil.save(os.path.join(save_dir, out_name))\n\nprint(\"✅ 所有 test crop 图像生成完毕，已保存至：\", save_dir)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T07:58:13.744129Z","iopub.execute_input":"2025-06-19T07:58:13.74435Z","iopub.status.idle":"2025-06-19T08:06:15.262995Z","shell.execute_reply.started":"2025-06-19T07:58:13.744329Z","shell.execute_reply":"2025-06-19T08:06:15.262272Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 分类模型结构与加载","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\n\nclass SimpleUNet(nn.Module):\n    def __init__(self, in_channels=3, out_channels=19):\n        super(SimpleUNet, self).__init__()\n\n        def conv_block(in_ch, out_ch):\n            return nn.Sequential(\n                nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_ch),\n                nn.LeakyReLU(inplace=True),\n                nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_ch),\n                nn.LeakyReLU(inplace=True)\n            )\n\n        self.enc1 = conv_block(in_channels, 64)\n        self.pool1 = nn.MaxPool2d(2)\n\n        self.enc2 = conv_block(64, 128)\n        self.pool2 = nn.MaxPool2d(2)\n\n        self.bottleneck = conv_block(128, 256)\n\n        self.up1 = nn.ConvTranspose2d(256, 128, 2, stride=2)\n        self.dec1 = conv_block(256, 128)\n\n        self.up2 = nn.ConvTranspose2d(128, 64, 2, stride=2)\n        self.dec2 = conv_block(128, 64)\n\n        self.final = nn.Conv2d(64, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        p1 = self.pool1(e1)\n\n        e2 = self.enc2(p1)\n        p2 = self.pool2(e2)\n\n        b = self.bottleneck(p2)\n\n        u1 = self.up1(b)\n        d1 = self.dec1(torch.cat([u1, e2], dim=1))\n\n        u2 = self.up2(d1)\n        d2 = self.dec2(torch.cat([u2, e1], dim=1))\n\n        out = self.final(d2)\n        return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:06:15.264139Z","iopub.execute_input":"2025-06-19T08:06:15.264371Z","iopub.status.idle":"2025-06-19T08:06:15.273195Z","shell.execute_reply.started":"2025-06-19T08:06:15.264347Z","shell.execute_reply":"2025-06-19T08:06:15.272322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 进行推理预测","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom tqdm import tqdm\nfrom PIL import Image\nimport torch\nimport torch.nn.functional as F\nfrom torchvision import transforms\n\n# ✅ 模型结构\nmodel = SimpleUNet(in_channels=3, out_channels=19)\nmodel_path = '/kaggle/input/unet_weak_supervision_epoch5_losweight/pytorch/default/1/unet_weak_supervision_epoch5_losweight.pth'\nmodel.load_state_dict(torch.load(model_path, map_location='cpu'))\nmodel.eval()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\nprint(f\"✅ 模型加载完成，使用设备: {device}\")\n\n# ✅ 预处理\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n])\n\n# ✅ crop 路径\ncrop_dir = '/kaggle/working/cell_crops_test_0.2'\nresults = []\n\n# ✅ tqdm 进度条遍历 crop 图像\nfor fname in tqdm(os.listdir(crop_dir), desc=\"🔍 正在推理\", unit=\"crop\"):\n    if not fname.endswith('.png'):\n        continue\n\n    try:\n        image_id = fname.split('_')[0]\n        class_id = int(fname.split('_')[1].replace('class', ''))\n        cell_id = int(fname.split('_')[2].replace('cell', '').replace('.png', ''))\n    except:\n        print(f\"[跳过] 文件名无法解析: {fname}\")\n        continue\n\n    img_path = os.path.join(crop_dir, fname)\n    img = Image.open(img_path).convert('RGB')\n    img_tensor = transform(img).unsqueeze(0).to(device)\n\n    with torch.no_grad():\n        output = model(img_tensor)  # [1, 19, H, W]\n        probs = torch.sigmoid(output).squeeze(0).cpu().numpy()  # [19, H, W]\n\n    # 每类的最大像素概率（也可以改成 mean）\n    conf_per_class = probs.reshape(19, -1).max(axis=1)\n    pred_classes = [i for i, p in enumerate(conf_per_class) if p > 0]\n\n    results.append({\n        'filename': fname,\n        'image_id': image_id,\n        'cell_id': cell_id,\n        'true_class': class_id,\n        'pred_classes': pred_classes,\n        'probs': [round(float(p), 4) for p in conf_per_class],\n    })\n\n# ✅ 保存 CSV\ndf = pd.DataFrame(results)\ndf.to_csv('/kaggle/working/unet_predictions_on_crops.csv', index=False)\nprint(\"✅ 已保存至 /kaggle/working/unet_predictions_on_crops.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:27:18.76625Z","iopub.execute_input":"2025-06-19T08:27:18.766581Z","iopub.status.idle":"2025-06-19T08:28:19.69841Z","shell.execute_reply.started":"2025-06-19T08:27:18.766552Z","shell.execute_reply":"2025-06-19T08:28:19.697772Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom tqdm import tqdm\nfrom PIL import Image\nimport torch\nimport torch.nn.functional as F\nfrom torchvision import transforms\n\n# ✅ 模型结构\nmodel = SimpleUNet(in_channels=3, out_channels=19)\nmodel_path = '/kaggle/input/unet_weak_supervision_epoch5_losweight/pytorch/default/1/unet_weak_supervision_epoch5_losweight.pth'\nmodel.load_state_dict(torch.load(model_path, map_location='cpu'))\nmodel.eval()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\nprint(f\"✅ 模型加载完成，使用设备: {device}\")\n\n# ✅ 预处理\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n])\n\n# ✅ crop 路径\ncrop_dir = '/kaggle/working/cell_crops_test_0.2'\nresults = []\n\n# ✅ tqdm 进度条遍历 crop 图像\nfor fname in tqdm(os.listdir(crop_dir), desc=\"🔍 正在推理\", unit=\"crop\"):\n    if not fname.endswith('.png'):\n        continue\n\n    try:\n        image_id = fname.split('_')[0]\n        class_id = int(fname.split('_')[1].replace('class', ''))\n        cell_id = int(fname.split('_')[2].replace('cell', '').replace('.png', ''))\n    except:\n        print(f\"[跳过] 文件名无法解析: {fname}\")\n        continue\n\n    img_path = os.path.join(crop_dir, fname)\n    img = Image.open(img_path).convert('RGB')\n    img_tensor = transform(img).unsqueeze(0).to(device)\n\n    with torch.no_grad():\n        output = model(img_tensor)  # [1, 19, H, W]\n        probs = torch.sigmoid(output).squeeze(0).cpu().numpy()  # [19, H, W]\n\n    # 每类的最大像素概率（无过滤）\n    conf_per_class = probs.reshape(19, -1).max(axis=1)\n\n    results.append({\n        'filename': fname,\n        'image_id': image_id,\n        'cell_id': cell_id,\n        'true_class': class_id,\n        'probs': [round(float(p), 4) for p in conf_per_class],\n    })\n\n# ✅ 保存完整 CSV（不过滤 pred_classes）\ndf = pd.DataFrame(results)\ndf.to_csv('/kaggle/working/unet_predictions_raw_probs.csv', index=False)\nprint(\"✅ 已保存未经过滤的预测结果至 /kaggle/working/unet_predictions_raw_probs.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:37:20.743368Z","iopub.execute_input":"2025-06-19T08:37:20.743761Z","iopub.status.idle":"2025-06-19T08:38:21.898513Z","shell.execute_reply.started":"2025-06-19T08:37:20.743729Z","shell.execute_reply":"2025-06-19T08:38:21.897618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\npd.read_csv('/kaggle/working/unet_predictions_raw_probs.csv').head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:38:25.44169Z","iopub.execute_input":"2025-06-19T08:38:25.442026Z","iopub.status.idle":"2025-06-19T08:38:25.465992Z","shell.execute_reply.started":"2025-06-19T08:38:25.441997Z","shell.execute_reply":"2025-06-19T08:38:25.465139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom collections import defaultdict\n\n# 加载 CSV\ndf = pd.read_csv('/kaggle/working/unet_predictions_raw_probs.csv')\n\n# 统计结构：每个 class 对应的 confidence 列表\nclass_confidences = defaultdict(list)\n\nfor _, row in df.iterrows():\n    # 将字符串形式的列表转换为真正的 Python list\n    probs = eval(row['probs'])  # [19]\n\n    for cls_id, prob in enumerate(probs):\n        class_confidences[cls_id].append(prob)\n\n# 分析每类的 confidence 分布\nprint(\"📊 每类预测的 confidence 分布：\")\nfor cls_id in sorted(class_confidences.keys()):\n    values = np.array(class_confidences[cls_id])\n    print(f\"Class {cls_id}: count={len(values)}, mean={values.mean():.4f}, std={values.std():.4f}, min={values.min():.4f}, max={values.max():.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:41:34.790364Z","iopub.execute_input":"2025-06-19T08:41:34.790707Z","iopub.status.idle":"2025-06-19T08:41:35.184577Z","shell.execute_reply.started":"2025-06-19T08:41:34.790679Z","shell.execute_reply":"2025-06-19T08:41:35.183654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(16, 12))\nfor i, cls in enumerate(sorted(class_confidences)):\n    plt.subplot(5, 4, i+1)\n    plt.hist(class_confidences[cls], bins=20, color='skyblue')\n    plt.title(f\"Class {cls} (n={len(class_confidences[cls])})\")\n    plt.xlabel(\"Confidence\")\n    plt.ylabel(\"Count\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:41:49.501373Z","iopub.execute_input":"2025-06-19T08:41:49.501688Z","iopub.status.idle":"2025-06-19T08:41:52.112185Z","shell.execute_reply.started":"2025-06-19T08:41:49.501663Z","shell.execute_reply":"2025-06-19T08:41:52.111473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter\nfrom ast import literal_eval\nimport pandas as pd\n\ndf = pd.read_csv('/kaggle/working/unet_predictions_on_crops.csv')\ndf['pred_classes'] = df['pred_classes'].apply(literal_eval)\nall_classes = [c for sub in df['pred_classes'] for c in sub]\nprint(Counter(all_classes))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-19T08:28:38.354902Z","iopub.execute_input":"2025-06-19T08:28:38.355184Z","iopub.status.idle":"2025-06-19T08:28:38.483777Z","shell.execute_reply.started":"2025-06-19T08:28:38.35516Z","shell.execute_reply":"2025-06-19T08:28:38.483032Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# fliter_1 confidence","metadata":{}}]}