{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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":12067245,"sourceType":"datasetVersion","datasetId":7595523},{"sourceId":12124945,"sourceType":"datasetVersion","datasetId":7634750},{"sourceId":244656261,"sourceType":"kernelVersion"},{"sourceId":251505556,"sourceType":"kernelVersion"},{"sourceId":486729,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":388152,"modelId":407157}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":4093.474804,"end_time":"2025-06-25T09:12:03.135862","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-06-25T08:03:49.661058","version":"2.2.2"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nimport os\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as nnf","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":1.972183,"end_time":"2025-06-25T08:03:56.334349","exception":false,"start_time":"2025-06-25T08:03:54.362166","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:57:51.534511Z","iopub.execute_input":"2025-07-24T02:57:51.534829Z","iopub.status.idle":"2025-07-24T02:58:01.755004Z","shell.execute_reply.started":"2025-07-24T02:57:51.534809Z","shell.execute_reply":"2025-07-24T02:58:01.754383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"PYDEVD_DISABLE_FILE_VALIDATION\"] = \"1\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:01.755548Z","iopub.execute_input":"2025-07-24T02:58:01.755865Z","iopub.status.idle":"2025-07-24T02:58:01.759259Z","shell.execute_reply.started":"2025-07-24T02:58:01.755848Z","shell.execute_reply":"2025-07-24T02:58:01.758748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#  设定参数","metadata":{"papermill":{"duration":0.012852,"end_time":"2025-06-25T08:03:56.361863","exception":false,"start_time":"2025-06-25T08:03:56.349011","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom torchvision import models\nfrom tqdm import tqdm\n\nNUM_CLASSES = 19\nBATCH_SIZE = 32\nEPOCHS = 3\nLR = 1e-4\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n","metadata":{"papermill":{"duration":0.067305,"end_time":"2025-06-25T08:03:56.442311","exception":false,"start_time":"2025-06-25T08:03:56.375006","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:01.760991Z","iopub.execute_input":"2025-07-24T02:58:01.761163Z","iopub.status.idle":"2025-07-24T02:58:01.833078Z","shell.execute_reply.started":"2025-07-24T02:58:01.761149Z","shell.execute_reply":"2025-07-24T02:58:01.832375Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 查看mask","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\n\n# 加载一个 npz 掩膜文件\npath = \"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/mask/train/cell/0042017c-bba4-11e8-b2b9-ac1f6b6435d0.npz\"\ndata = np.load(path, allow_pickle=True)\nmask = data[\"arr_0\"]  # 或 data.files[0]，已经确认是 'arr_0'\n\n# 打印掩膜信息\nprint(\"掩膜 shape:\", mask.shape)\nprint(\"掩膜类型:\", type(mask))\nprint(\"掩膜像素取值种类:\", np.unique(mask))\n\n# 可视化\nplt.imshow(mask, cmap=\"tab20\")\nplt.colorbar()\nplt.title(\"Structure Mask Visualization\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:01.833938Z","iopub.execute_input":"2025-07-24T02:58:01.834213Z","iopub.status.idle":"2025-07-24T02:58:02.510755Z","shell.execute_reply.started":"2025-07-24T02:58:01.834184Z","shell.execute_reply":"2025-07-24T02:58:02.5099Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 生成weak supervision的csv(带label，两者结合)","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# 路径\nbbox_path = \"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/train_bbox_filtered.csv\"\ntrain_path = \"/kaggle/input/hpa-single-cell-image-classification/train.csv\"\n\n# 读取数据\nbbox_df = pd.read_csv(bbox_path)\ntrain_df = pd.read_csv(train_path)\n\n# 预处理标签字段\ndef parse_label(label_str):\n    return [int(x) for x in label_str.split('|') if x != '']\n\ntrain_df['label_list'] = train_df['Label'].map(parse_label)\ntrain_df = train_df.rename(columns={\"ID\": \"image_id\"})\n\n# 构建映射（image_id → multi-hot 向量）\nimageid_to_vector = {}\nfor _, row in train_df.iterrows():\n    vec = [0] * 19\n    for i in row['label_list']:\n        if 0 <= i < 19:\n            vec[i] = 1\n    imageid_to_vector[row['image_id']] = vec\n\n# 添加 label_vector 列\nbbox_df['label_vector'] = bbox_df['image_id'].map(imageid_to_vector)\n\n# 保存结果\nbbox_df.to_csv(\"/kaggle/working/bbox_with_label_vector.csv\", index=False)\n\nprint(\"保存成功！共 %d 个细胞，字段名如下：\" % len(bbox_df))\nprint(bbox_df.columns.tolist())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:02.511413Z","iopub.execute_input":"2025-07-24T02:58:02.511636Z","iopub.status.idle":"2025-07-24T02:58:06.743012Z","shell.execute_reply.started":"2025-07-24T02:58:02.511593Z","shell.execute_reply":"2025-07-24T02:58:06.742322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 看一下npz文件里都有什么并创建一个目录¶","metadata":{}},{"cell_type":"code","source":"# ========= 可选 DEBUG：查看单个 .npz 掩膜结构 =========\nimport numpy as np\nDEBUG = False\nif DEBUG:\n    path = \"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/mask/train/cell/0042017c-bba4-11e8-b2b9-ac1f6b6435d0.npz\"\n    data = np.load(path, allow_pickle=True)\n    content = data[\"arr_0\"]\n    print(\"字段类型：\", type(content))\n    if isinstance(content, np.ndarray):\n        print(\"shape:\", content.shape)\n        print(\"取值种类:\", np.unique(content))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:06.743775Z","iopub.execute_input":"2025-07-24T02:58:06.743998Z","iopub.status.idle":"2025-07-24T02:58:06.748383Z","shell.execute_reply.started":"2025-07-24T02:58:06.743972Z","shell.execute_reply":"2025-07-24T02:58:06.747862Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 读取RYB图像，也就是RGB，然后resize","metadata":{}},{"cell_type":"code","source":"from PIL import Image  # 确保已导入\n\nclass WeakSupervisionDataset(Dataset):\n    def __init__(self, csv_path, image_dir, mask_dir, transform=None, crop_size=512, target_size=2048):\n        self.df = pd.read_csv(csv_path)\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.transform = transform\n        self.crop_size = crop_size\n        self.target_w = target_size\n        self.target_h = target_size\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row['image_id']\n        cell_id = row['cell_id']\n        x0, y0, x1, y1 = int(row['x0']), int(row['y0']), int(row['x1']), int(row['y1'])\n\n        # === Step 1: 读取 R/Y/B 图像，并记录原始尺寸 ===\n        r = cv2.imread(os.path.join(self.image_dir, f\"{image_id}_red.png\"), cv2.IMREAD_GRAYSCALE)\n        y = cv2.imread(os.path.join(self.image_dir, f\"{image_id}_yellow.png\"), cv2.IMREAD_GRAYSCALE)\n        b = cv2.imread(os.path.join(self.image_dir, f\"{image_id}_blue.png\"), cv2.IMREAD_GRAYSCALE)\n\n        original_h, original_w = r.shape[:2]  # ⚠️ 原图尺寸\n\n        # === Step 2: resize 图像，并构造多通道 ===\n        r = cv2.resize(r, (self.target_w, self.target_h))\n        y = cv2.resize(y, (self.target_w, self.target_h))\n        b = cv2.resize(b, (self.target_w, self.target_h))\n        image = np.stack([r, y, b], axis=-1)  # [2048, 2048, 3]\n\n        # === Step 3: 读取并 resize 掩膜 ===\n        npz_path = os.path.join(self.mask_dir, f\"{image_id}.npz\")\n        mask = np.load(npz_path, allow_pickle=True)[\"arr_0\"]\n        mask = cv2.resize(mask, (self.target_w, self.target_h), interpolation=cv2.INTER_NEAREST)\n\n        # === Step 4: ⚠️ 同步缩放 bbox 坐标 ===\n        scale_x = self.target_w / original_w\n        scale_y = self.target_h / original_h\n        x0 = int(x0 * scale_x)\n        x1 = int(x1 * scale_x)\n        y0 = int(y0 * scale_y)\n        y1 = int(y1 * scale_y)\n\n        # === Step 5: 从 mask 中定位该 cell 的 tight bbox ===\n        target_mask = (mask == int(cell_id)).astype(np.uint8)  # binary mask\n        ys, xs = np.where(target_mask)\n\n        if len(xs) == 0 or len(ys) == 0:\n            # === Step 6: 异常处理（目标 cell 不存在） ===\n            dummy = np.zeros((self.crop_size, self.crop_size, 3), dtype=np.uint8)\n            crop_img = Image.fromarray(dummy)\n            label_vector = eval(row['label_vector'])\n            crop_tensor = transforms.ToTensor()(crop_img)\n            label_tensor = torch.tensor(label_vector).float()\n            mask_tensor = torch.zeros((1, self.crop_size, self.crop_size))\n            return crop_tensor, mask_tensor, label_tensor, image_id, cell_id\n\n        # === Step 6: 基于 mask 定 tight bbox 并加 padding ===\n        x_min, x_max = xs.min(), xs.max()\n        y_min, y_max = ys.min(), ys.max()\n        pad = 10\n\n        cx = (x_min + x_max) // 2\n        cy = (y_min + y_max) // 2\n        half_size = max(x_max - x_min, y_max - y_min) // 2 + pad\n\n        new_x0 = max(0, cx - half_size)\n        new_x1 = min(self.target_w, cx + half_size)\n        new_y0 = max(0, cy - half_size)\n        new_y1 = min(self.target_h, cy + half_size)\n\n        # === Step 7: crop 图像、掩膜，并 resize ===\n        crop_img = image[new_y0:new_y1, new_x0:new_x1, :]  # [H,W,3]\n        crop_mask = target_mask[new_y0:new_y1, new_x0:new_x1]  # [H,W]\n\n        crop_img = cv2.resize(crop_img, (self.crop_size, self.crop_size))\n        crop_mask = cv2.resize(crop_mask, (self.crop_size, self.crop_size), interpolation=cv2.INTER_NEAREST)\n\n        # 淡化背景（非目标 cell 的像素）\n        crop_img[crop_mask == 0] = (crop_img[crop_mask == 0] * 0.1).astype(np.uint8)\n\n        crop_img = Image.fromarray(crop_img.astype(np.uint8))  # PIL Image\n\n        # === Step 8: 图像 transform / 归一化 ===\n        if self.transform:\n            augmented = self.transform(image=np.array(crop_img), mask=crop_mask)\n            crop_tensor = augmented[\"image\"]\n            mask_tensor = augmented[\"mask\"].unsqueeze(0).float()  # [1, H, W]\n        else:\n            crop_np = np.array(crop_img).astype(np.float32) / 255.0\n            crop_np = np.transpose(crop_np, (2, 0, 1))  # [C,H,W]\n            crop_tensor = torch.tensor(crop_np).float()\n            mask_tensor = torch.tensor(crop_mask).unsqueeze(0).float()\n\n        # === Step 9: 标签处理 ===\n        label_vector = eval(row['label_vector'])\n        label_tensor = torch.tensor(label_vector).float()\n\n        return crop_tensor, mask_tensor, label_tensor, image_id, cell_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:06.749203Z","iopub.execute_input":"2025-07-24T02:58:06.749425Z","iopub.status.idle":"2025-07-24T02:58:06.767442Z","shell.execute_reply.started":"2025-07-24T02:58:06.74941Z","shell.execute_reply":"2025-07-24T02:58:06.766745Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"所以说明在进行真正的weak supervision之前，这个weak lable完全不可能删除。因为这是赋值的准备。比如说这个16|2|0|13。我把这四个class，都给到了这个image id的所有cell里面。这就导致重复的filename一定会出现四回，且不可能被删除\n\n","metadata":{"papermill":{"duration":0.014332,"end_time":"2025-06-25T08:03:57.720076","exception":false,"start_time":"2025-06-25T08:03:57.705744","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# define PyTorch Dataset","metadata":{"papermill":{"duration":0.015337,"end_time":"2025-06-25T08:04:05.170681","exception":false,"start_time":"2025-06-25T08:04:05.155344","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"在这里想改变一下dataframe的构成。一次性构造 multi-hot 标签向量，而不是每个 class 一行 模型在一次训练中学到“它是 class 0”，下一次又学到“它是 class 2”，但每次都把其它 class（如 2 或 0）当成负类\n\n","metadata":{"papermill":{"duration":0.015017,"end_time":"2025-06-25T08:04:05.200776","exception":false,"start_time":"2025-06-25T08:04:05.185759","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from PIL import Image\nfrom torch.utils.data import DataLoader\nimport pandas as pd\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport numpy as np\nimport ast\nfrom sklearn.model_selection import train_test_split\n\n\nDEBUG = True  # ✅ 是否打印一张样本图像和标签信息\n\n# 读取已构建的 cell-level label 向量\ncsv_path = \"/kaggle/working/bbox_with_label_vector.csv\"\ndf_valid = pd.read_csv(csv_path)\n\n# 定义 transform\ntransform = A.Compose([\n    A.Resize(256, 256),  # 可选改为 512\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),\n    ToTensorV2()\n])\n\n# 构建 Dataset\ndataset = WeakSupervisionDataset(\n    csv_path=csv_path,\n    image_dir=\"/kaggle/input/hpa-single-cell-image-classification/train\",\n    mask_dir=\"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/mask/train/cell\",\n    transform=transform,\n    target_size=2048\n)\n\n# 构建 DataLoader\n# ✅ 划分训练集和验证集索引\ntrain_idx, val_idx = train_test_split(\n    np.arange(len(dataset)),\n    test_size=0.01,  # 验证集比例\n    random_state=42,\n    shuffle=True\n)\n\n# ✅ 创建子集数据集\ntrain_dataset = torch.utils.data.Subset(dataset, train_idx)\nval_dataset = torch.utils.data.Subset(dataset, val_idx)\n\n# ✅ 分别构建 DataLoader\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=32,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=32,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n\n\n# ✅ 打印一张样本信息（可选）\nif DEBUG:\n    crop_tensor, mask_tensor, label_tensor, image_id, cell_id = dataset[0]\n    print(f\"图像大小: {crop_tensor.shape}\")\n    print(f\"掩膜大小: {mask_tensor.shape}\")\n    print(f\"标签向量: {label_tensor}\")\n    print(f\"图像 ID: {image_id}, Cell ID: {cell_id}\")\n\n# ✅ 每类样本数统计（图像级标签分配到细胞）\nlabel_matrix = df_valid[\"label_vector\"].apply(lambda x: np.array(ast.literal_eval(x)))\nlabel_sum = np.sum(np.stack(label_matrix.values), axis=0)\nprint(\"每个类别的样本数:\", label_sum.astype(int))\n","metadata":{"papermill":{"duration":0.221942,"end_time":"2025-06-25T08:04:05.475359","exception":false,"start_time":"2025-06-25T08:04:05.253417","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:06.768261Z","iopub.execute_input":"2025-07-24T02:58:06.768512Z","iopub.status.idle":"2025-07-24T02:58:12.950586Z","shell.execute_reply.started":"2025-07-24T02:58:06.768495Z","shell.execute_reply":"2025-07-24T02:58:12.949936Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"224 x 224 是使用 transforms.Resize((224, 224)) 指定的尺寸;[3] 是 RGB 通道数；表示该 cell crop 属于第 15 个类别（索引从 0 开始）","metadata":{"papermill":{"duration":0.015708,"end_time":"2025-06-25T08:04:05.942314","exception":false,"start_time":"2025-06-25T08:04:05.926606","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 看一下截取的mask 还有crop","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nimport random\nimport cv2\nimport numpy as np\n\nVISUALIZE = True  # ✅ 控制是否执行可视化代码块\n\n# === 路径设置 ===\ncsv_path = \"/kaggle/working/bbox_with_label_vector.csv\"\nimage_dir = \"/kaggle/input/hpa-single-cell-image-classification/train\"\nmask_dir = \"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/mask/train/cell\"\n\n# === 图像增强 transform ===\ntransform = A.Compose([\n    A.Resize(256, 256),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),\n    ToTensorV2()\n])\n\n# === 构建训练集 dataset ===\ntrain_dataset = WeakSupervisionDataset(\n    csv_path=csv_path,\n    image_dir=image_dir,\n    mask_dir=mask_dir,\n    transform=transform,\n    target_size=2048\n)\n\n# === 可视化函数 1：随机展示 crop 图像 ===\ndef visualize_random_crops(dataset, num_samples=16, rows=4, cols=4, seed=42):\n    random.seed(seed)\n    indices = random.sample(range(len(dataset)), num_samples)\n\n    fig, axes = plt.subplots(rows, cols, figsize=(4 * cols, 4 * rows))\n    axes = axes.flatten()\n\n    for ax, idx in zip(axes, indices):\n        crop_tensor, mask_tensor, label_tensor, image_id, cell_id = dataset[idx]\n        img_np = crop_tensor.permute(1, 2, 0).cpu().numpy()\n        img_np = (img_np * 0.5 + 0.5).clip(0, 1)\n        ax.imshow(img_np)\n        ax.set_title(f\"Image: {image_id[-5:]}\\nCell: {cell_id}\", fontsize=8)\n        ax.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# === 可视化函数 2：原图 + mask + crop ===\ndef visualize_full_and_crop(dataset, num_samples=5, seed=42):\n    np.random.seed(seed)\n    indices = np.random.choice(len(dataset), num_samples, replace=False)\n\n    for idx in indices:\n        image_tensor, mask_tensor, label_vector, image_id, cell_id = dataset[idx]\n        crop_img = image_tensor.permute(1, 2, 0).numpy()\n        crop_img = (crop_img * 0.5 + 0.5).clip(0, 1)\n        mask_np = mask_tensor.squeeze().numpy()\n\n        # 原图 resize\n        r = cv2.imread(f\"{image_dir}/{image_id}_red.png\", cv2.IMREAD_GRAYSCALE)\n        y = cv2.imread(f\"{image_dir}/{image_id}_yellow.png\", cv2.IMREAD_GRAYSCALE)\n        b = cv2.imread(f\"{image_dir}/{image_id}_blue.png\", cv2.IMREAD_GRAYSCALE)\n        rgb = np.stack([r, y, b], axis=-1)\n        rgb = cv2.resize(rgb, (2048, 2048)).astype(np.float32) / 255.0\n\n        # 掩膜 overlay\n        full_mask = np.load(f\"{mask_dir}/{image_id}.npz\", allow_pickle=True)[\"arr_0\"]\n        full_mask = cv2.resize(full_mask, (2048, 2048), interpolation=cv2.INTER_NEAREST)\n        cell_mask = (full_mask == cell_id).astype(np.uint8)\n        overlay = rgb.copy()\n        overlay[cell_mask == 1] = [1.0, 0.0, 0.0]\n\n        # 显示\n        plt.figure(figsize=(10, 4))\n        plt.subplot(1, 2, 1)\n        plt.imshow(overlay)\n        plt.title(f\"Full Image + Mask\\n{image_id[-5:]} | Cell {cell_id}\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 2, 2)\n        plt.imshow(crop_img)\n        plt.title(\"Cropped Cell Image\")\n        plt.axis(\"off\")\n        plt.tight_layout()\n        plt.show()\n\n# === 是否执行可视化 ===\nif VISUALIZE:\n    visualize_random_crops(train_dataset)\n    visualize_full_and_crop(train_dataset)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:12.952938Z","iopub.execute_input":"2025-07-24T02:58:12.953186Z","iopub.status.idle":"2025-07-24T02:58:27.685102Z","shell.execute_reply.started":"2025-07-24T02:58:12.953167Z","shell.execute_reply":"2025-07-24T02:58:27.684316Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"在这里我们使用了旋转，让图像可以得到更多的训练","metadata":{}},{"cell_type":"markdown","source":"# los_weight # define model and loss function and optimizer and loop","metadata":{"papermill":{"duration":0.015411,"end_time":"2025-06-25T08:04:05.97321","exception":false,"start_time":"2025-06-25T08:04:05.957799","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nimport numpy as np\nimport torch.nn as nn\nfrom torchvision import models\nfrom tqdm import tqdm\n\n# ========= 参数配置 =========\nNUM_CLASSES = 19\nLR = 1e-4\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmeta_path = '/kaggle/input/hpa-single-cell-image-classification/train.csv'\nweight_path = '/kaggle/working/pos_weight.npy'\nbest_model_path = \"/kaggle/working/best_model_train.pth\"\nEPOCHS = 3\nTRAIN = True\n\n# ========= 计算 pos_weight =========\ndef get_pos_weight(csv_path: str, quiet: bool = True):\n    df_meta = pd.read_csv(csv_path)\n    df_meta['Label_list'] = df_meta['Label'].apply(lambda x: list(map(int, x.split('|'))))\n\n    label_counts = np.zeros(NUM_CLASSES, dtype=int)\n    for labels in df_meta['Label_list']:\n        for cls in labels:\n            label_counts[cls] += 1\n\n    class_freq = label_counts / label_counts.sum()\n    pos_weight = np.log((1 - class_freq) / class_freq + 1e-8)\n\n    pos_weight[18] = 0.0\n    pos_weight[0] = min(pos_weight[0], 3.5)\n    pos_weight[16] = min(pos_weight[16], 3)\n\n    if not quiet:\n        print(\"每个类别的出现次数:\", label_counts.tolist())\n        for i in range(len(pos_weight)):\n            print(f\"class {i:2d}: pos_weight = {pos_weight[i]:.4f}\")\n    return pos_weight, label_counts\n\n# ========= 加载 pos_weight =========\nif \"pos_weight\" not in globals():\n    if os.path.exists(weight_path):\n        pos_weight = np.load(weight_path)\n        print(f\"✅ pos_weight 已加载自: {weight_path}\")\n    else:\n        pos_weight, label_counts = get_pos_weight(meta_path, quiet=False)\n        np.save(weight_path, pos_weight)\n        print(f\"✅ 重新计算并保存 pos_weight 至: {weight_path}\")\n\n\nfrom torchvision.models import ResNet18_Weights\n\n# ========= 构建模型（使用预训练） =========\nmodel = models.resnet18(weights=None)  # 不要重新加载预训练\n\n# ========= 替换 classifier 层 =========\nmodel.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)  # 替换 classifier 层\nmodel.load_state_dict(torch.load(\"/kaggle/input/deforz-layer-3-4-resnet18-training/pytorch/default/1/best_model_train.pth\", map_location=DEVICE))\n\n# ========= 解冻layer1 ,2,  layer3、layer4 和 fc（推荐做法） =========\nfor name, param in model.named_parameters():\n    if any(x in name for x in [\"layer1\", \"layer2\", \"layer3\", \"layer4\", \"fc\"]):\n        param.requires_grad = True\n    else:\n        param.requires_grad = False\n\nmodel = model.to(DEVICE)\n\n# ========= Loss 和 Optimizer =========\npos_weight_tensor = torch.tensor(pos_weight, dtype=torch.float32).to(DEVICE)\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)\n\n# ✅ 只优化需要梯度的部分（不会训练冻结层）\noptimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=LR)\n\n# ========= 训练模型 =========\nif TRAIN and not os.path.exists(best_model_path):\n    print(\"🟢 开始训练模型...\")\n    best_val_loss = float('inf')\n    for epoch in range(EPOCHS):\n        model.train()\n        total_loss = 0.0\n        step = 0\n\n        train_loader_tqdm = tqdm(train_loader, desc=f\"[Train] Epoch {epoch+1}\")\n        for images, _, labels, image_ids, cell_ids in train_loader_tqdm:\n            images = images.to(DEVICE)\n            labels = labels.to(DEVICE)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n            train_loader_tqdm.set_postfix(loss=loss.item())\n\n            step += 1\n            #if step > 40:  # 若调试时可取消该限制\n                #break\n\n        avg_train_loss = total_loss / step\n        print(f\"✅ Epoch {epoch+1}/{EPOCHS} 完成，平均训练损失: {avg_train_loss:.4f}\")\n\n        # ========= 验证（可选） =========\n        if 'val_loader' in globals():\n            model.eval()\n            val_loss = 0.0\n            val_step = 0\n            with torch.no_grad():\n                for images, _, labels, *_ in val_loader:\n                    images = images.to(DEVICE)\n                    labels = labels.to(DEVICE)\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    val_step += 1\n                    if val_step >= 40:  # 🟢 控制验证只跑前 40 个 batch\n                        break\n\n            avg_val_loss = val_loss / val_step\n            print(f\"🧪 验证损失: {avg_val_loss:.4f}\")\n\n            if avg_val_loss < best_val_loss:\n                best_val_loss = avg_val_loss\n                torch.save(model.state_dict(), best_model_path)\n                print(f\"✅ 新最佳模型（val）保存至: {best_model_path}\")\n        else:\n            if avg_train_loss < best_val_loss:\n                best_val_loss = avg_train_loss\n                torch.save(model.state_dict(), best_model_path)\n                print(f\"✅ 新最佳模型（train）保存至: {best_model_path}\")\nelse:\n    print(f\"⚠️ 已存在模型 {best_model_path}，跳过训练阶段。\")\n\nprint(\"✅✅✅ 训练真正完成于此处 ✅✅✅\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T02:58:27.68605Z","iopub.execute_input":"2025-07-24T02:58:27.686268Z","iopub.status.idle":"2025-07-24T03:04:02.170403Z","shell.execute_reply.started":"2025-07-24T02:58:27.686251Z","shell.execute_reply":"2025-07-24T03:04:02.169348Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"我们不在这里进行csv的保存，因为模型对每一个cell都进行19个vector的输出。如果在这里保存会出现非常大的误差\n\n","metadata":{"papermill":{"duration":3.696052,"end_time":"2025-06-25T09:00:32.275185","exception":false,"start_time":"2025-06-25T09:00:28.579133","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# threshold and save as csv","metadata":{"papermill":{"duration":3.686182,"end_time":"2025-06-25T09:00:39.614417","exception":false,"start_time":"2025-06-25T09:00:35.928235","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pandas as pd\nimport torch\nfrom tqdm import tqdm\nimport os\n\nPREDICT = True  # ✅ 控制是否执行推理\nmodel_path = \"/kaggle/working/best_model_train.pth\"\npred_csv_path = \"/kaggle/working/cell_level_predictions_conf_05_train.csv\"\nthreshold = 0.5\nNUM_CLASSES = 19\nmax_steps = 30\n\nif PREDICT:\n    # 如果文件存在就不重复推理\n    if os.path.exists(pred_csv_path):\n        print(f\"⚠️ 已存在预测结果文件，跳过预测: {pred_csv_path}\")\n    else:\n        # 加载模型\n        model.load_state_dict(torch.load(model_path, map_location=DEVICE))\n        model.eval()\n\n        results = []\n\n        with torch.no_grad():\n            for step, (images, _, labels, image_ids, cell_ids) in enumerate(\n                tqdm(train_loader, desc=\"🚀 Collecting predictions\")\n            ):\n                if step >= max_steps:\n                    break\n\n                images = images.to(DEVICE)\n                outputs = model(images)\n                probs = torch.sigmoid(outputs).cpu()\n\n                for i in range(images.size(0)):\n                    for class_id in range(NUM_CLASSES):\n                        conf = probs[i, class_id].item()\n                        if conf > threshold:\n                            results.append({\n                                \"image_id\": image_ids[i],\n                                \"cell_id\": int(cell_ids[i]),\n                                \"class_id\": class_id,\n                                \"confidence\": conf,\n                                \"filename\": f\"{image_ids[i]}_class{class_id}_cell{cell_ids[i]}.png\"\n                            })\n\n        # 保存结果\n        df_result = pd.DataFrame(results)\n        df_result.to_csv(pred_csv_path, index=False)\n        print(f\"✅ 推理完成，保存结果至: {pred_csv_path}，共 {len(df_result)} 条\")\n","metadata":{"papermill":{"duration":566.998145,"end_time":"2025-06-25T09:10:10.242683","exception":false,"start_time":"2025-06-25T09:00:43.244538","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:08:46.637906Z","iopub.execute_input":"2025-07-24T03:08:46.638228Z","iopub.status.idle":"2025-07-24T03:10:11.62804Z","shell.execute_reply.started":"2025-07-24T03:08:46.638208Z","shell.execute_reply":"2025-07-24T03:10:11.627249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取 CSV 文件\ndf = pd.read_csv('/kaggle/working/cell_level_predictions_conf_05_train.csv')\n\n# 显示前 15 行\nprint(df.head(15))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:19.411837Z","iopub.execute_input":"2025-07-24T03:10:19.41296Z","iopub.status.idle":"2025-07-24T03:10:19.42547Z","shell.execute_reply.started":"2025-07-24T03:10:19.412914Z","shell.execute_reply":"2025-07-24T03:10:19.424801Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# top 4","metadata":{"papermill":{"duration":3.973295,"end_time":"2025-06-25T09:10:18.207166","exception":false,"start_time":"2025-06-25T09:10:14.233871","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pandas as pd\nimport os\n\nFILTER_TOPK = True  # ✅ 控制是否执行 top4 筛选\n\ninput_path = \"/kaggle/working/cell_level_predictions_conf_05_train.csv\"\noutput_path = \"/kaggle/working/cell_level_predictions_conf_05_top4_train.csv\"\n\nif FILTER_TOPK:\n    if os.path.exists(output_path):\n        print(f\" 已存在 top4 筛选结果，跳过处理: {output_path}\")\n    else:\n        # 读取预测结果\n        df = pd.read_csv(input_path)\n\n        # 按 image_id 分组保留 top4 class（按 confidence）\n        def keep_top4_classes(group):\n            if len(group['class_id'].unique()) > 4:\n                return group.sort_values('confidence', ascending=False).drop_duplicates('class_id').head(4)\n            else:\n                return group\n\n        df_top4 = df.groupby('image_id', group_keys=False).apply(keep_top4_classes).reset_index(drop=True)\n\n        # 保存\n        df_top4.to_csv(output_path, index=False)\n        print(f\" 按 image_id 保留 top4 类别成功，保存到: {output_path}\")\n","metadata":{"papermill":{"duration":6.393525,"end_time":"2025-06-25T09:10:28.682004","exception":false,"start_time":"2025-06-25T09:10:22.288479","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:19.426643Z","iopub.execute_input":"2025-07-24T03:10:19.426859Z","iopub.status.idle":"2025-07-24T03:10:19.508691Z","shell.execute_reply.started":"2025-07-24T03:10:19.426842Z","shell.execute_reply":"2025-07-24T03:10:19.507986Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# match accuracy","metadata":{"papermill":{"duration":3.978757,"end_time":"2025-06-25T09:10:36.642202","exception":false,"start_time":"2025-06-25T09:10:32.663445","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"注意这个只是依靠我们前面run出来的image 进行推测，不是全部的提取","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom collections import defaultdict\n\n# ✅ 加载预测结果（top4 筛选后）\npred_df = pd.read_csv(\"/kaggle/working/cell_level_predictions_conf_05_top4_train.csv\")\n\n# ✅ 加载真实标签（你清洗过的 multi-hot label vector）\ntrue_df = pd.read_csv(\"/kaggle/working/bbox_with_label_vector.csv\")\n\n# ✅ 构建 pred_dict：[(image_id, cell_id) → set(class_id)]\npred_dict = defaultdict(set)\nfor _, row in pred_df.iterrows():\n    key = (row[\"image_id\"], row[\"cell_id\"])\n    pred_dict[key].add(int(row[\"class_id\"]))\n\n# ✅ 构建 true_dict：[(image_id, cell_id) → set(class_id)]\ntrue_dict = defaultdict(set)\nfor _, row in true_df.iterrows():\n    key = (row[\"image_id\"], row[\"cell_id\"])\n    label_vec = eval(row[\"label_vector\"])  # 转换为 list\n    for i, v in enumerate(label_vec):\n        if v == 1:\n            true_dict[key].add(i)\n\n# ✅ 聚合 cell-level 预测和真实\nmerged = defaultdict(lambda: {\"true\": set(), \"pred\": set()})\nall_keys = set(pred_dict.keys()) | set(true_dict.keys())\nfor key in all_keys:\n    merged[key]['true'] = true_dict.get(key, set())\n    merged[key]['pred'] = pred_dict.get(key, set())\n\n# ✅ 聚合到 image-level，只保留预测中出现的 image\npred_image_ids = set(pred_df['image_id'].unique())\nimage_summary = defaultdict(lambda: {\"true\": set(), \"pred\": set()})\nfor (image_id, cell_id), group in merged.items():\n    if image_id in pred_image_ids:  # ✅ 只统计做过预测的图像\n        image_summary[image_id]['true'].update(group['true'])\n        image_summary[image_id]['pred'].update(group['pred'])\n\n# ✅ 打印前15张图像的匹配信息\nprint(\"每张图像的预测匹配统计（仅包含做过预测的图像）:\\n\")\nimage_match_total = 0\nimage_true_total = 0\n\nfor i, (image_id, group) in enumerate(image_summary.items()):\n    true_set = group['true']\n    pred_set = group['pred']\n    matched = true_set & pred_set\n\n    if i < 15:\n        print(f\"Image {i+1}: {image_id}\")\n        print(f\"True class_id(s): {sorted(true_set)}\")\n        print(f\"Pred class_id(s): {sorted(pred_set)}\")\n        print(f\"Matched class_id(s): {sorted(matched)}\")\n        print(f\"Match count: {len(matched)} / {len(true_set)}\")\n        print(\"-\" * 50)\n\n    image_match_total += len(matched)\n    image_true_total += len(true_set)\n\n# ✅ 最终准确率\nif image_true_total > 0:\n    print(f\"\\n按预测出现的 image_id 统计的总体 Label-match Accuracy: {image_match_total / image_true_total:.4f}\")\nelse:\n    print(\"没有有效的真实标签用于匹配。\")\n","metadata":{"papermill":{"duration":19.932652,"end_time":"2025-06-25T09:11:00.578146","exception":false,"start_time":"2025-06-25T09:10:40.645494","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:19.509348Z","iopub.execute_input":"2025-07-24T03:10:19.509549Z","iopub.status.idle":"2025-07-24T03:10:26.23869Z","shell.execute_reply.started":"2025-07-24T03:10:19.509532Z","shell.execute_reply":"2025-07-24T03:10:26.237982Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" weak supervision 模型（ResNet18）已经成功收敛，任务本身是 multi-label 多分类，每个 cell 实际上可能属于多个 class。我们在这里强行简化为单标签分类,所以准确率天然会偏低，说明模型基本已经学会根据 cell crop 图像判断其弱标签","metadata":{"papermill":{"duration":4.004783,"end_time":"2025-06-25T09:11:08.572961","exception":false,"start_time":"2025-06-25T09:11:04.568178","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"这个是我们在b5里面真正要用的，因为他们的confidence都很高，可能是因为泛化导致的，保留更有准确性","metadata":{"papermill":{"duration":3.993232,"end_time":"2025-06-25T09:11:16.558424","exception":false,"start_time":"2025-06-25T09:11:12.565192","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# what is in the csv","metadata":{"papermill":{"duration":3.964319,"end_time":"2025-06-25T09:11:24.516726","exception":false,"start_time":"2025-06-25T09:11:20.552407","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pandas as pd\n\n# 读取 CSV 文件\ndf = pd.read_csv('/kaggle/working/cell_level_predictions_conf_05_top4_train.csv')\n\n# 计算唯一 filename 数量\nnum_unique_filenames = df['filename'].nunique()\nnum_unique_images = df['image_id'].nunique()\nprint(f\"唯一 filename 数量为: {num_unique_filenames}\")\nprint(f\"唯一 image_id 数量为: {num_unique_images}\")","metadata":{"papermill":{"duration":4.114716,"end_time":"2025-06-25T09:11:32.724108","exception":false,"start_time":"2025-06-25T09:11:28.609392","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:26.240478Z","iopub.execute_input":"2025-07-24T03:10:26.240977Z","iopub.status.idle":"2025-07-24T03:10:26.250584Z","shell.execute_reply.started":"2025-07-24T03:10:26.240954Z","shell.execute_reply":"2025-07-24T03:10:26.249783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.read_csv('/kaggle/working/bbox_with_label_vector.csv')\n\nunique_image_ids = df['image_id'].nunique()\nprint(\"唯一的 image_id 数量:\", unique_image_ids)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:26.251317Z","iopub.execute_input":"2025-07-24T03:10:26.251524Z","iopub.status.idle":"2025-07-24T03:10:26.936419Z","shell.execute_reply.started":"2025-07-24T03:10:26.251508Z","shell.execute_reply":"2025-07-24T03:10:26.935635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n# 读取 CSV 文件\ndf = pd.read_csv('/kaggle/working/cell_level_predictions_conf_05_top4_train.csv')\n\n# 每个 image_id 下的唯一 pred_class 种类数量\nclass_counts_per_image = df.groupby('image_id')['class_id'].nunique()\n\n# 统计：有 N 种类的 image 有多少张（N=1, 2, ..., 10）\ncount_distribution = class_counts_per_image.value_counts().sort_index()\n\n# 绘图\nplt.figure(figsize=(8, 5))\ncount_distribution.plot(kind='bar')\nplt.title('Number of Unique Predicted Classes per Image')\nplt.xlabel('Number of Unique Predicted Classes')\nplt.ylabel('Number of Images')\nplt.xticks(rotation=0)\nplt.tight_layout()\nplt.show()","metadata":{"papermill":{"duration":4.33048,"end_time":"2025-06-25T09:11:41.052172","exception":false,"start_time":"2025-06-25T09:11:36.721692","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:26.937266Z","iopub.execute_input":"2025-07-24T03:10:26.937801Z","iopub.status.idle":"2025-07-24T03:10:27.108012Z","shell.execute_reply.started":"2025-07-24T03:10:26.937776Z","shell.execute_reply":"2025-07-24T03:10:27.107237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取 CSV 文件\ndf = pd.read_csv('/kaggle/working/cell_level_predictions_conf_05_top4_train.csv')\n\n# 按 class_id 分组，统计每类中唯一 cell_id 的数量\ncell_count_per_class = df.groupby('class_id')['cell_id'].nunique().sort_index()\n# 按 class_id 分组，统计每类出现在多少张不同图像中（以 filename 作为 proxy）\nfile_count_per_class = df.groupby('class_id')['filename'].nunique().sort_index()\n\n# 展示结果\nprint(\"每个 class 被预测为正的 cell 数量：\")\nprint(cell_count_per_class)\nprint(\"每个 class 在多少张不同图像中被 confident 地预测为正：\")\nprint(file_count_per_class)","metadata":{"papermill":{"duration":4.120506,"end_time":"2025-06-25T09:11:49.132146","exception":false,"start_time":"2025-06-25T09:11:45.01164","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:27.108888Z","iopub.execute_input":"2025-07-24T03:10:27.109153Z","iopub.status.idle":"2025-07-24T03:10:27.121025Z","shell.execute_reply.started":"2025-07-24T03:10:27.109131Z","shell.execute_reply":"2025-07-24T03:10:27.120231Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"as we do not have cell pred class larger than 5. just write this top k for double check","metadata":{"papermill":{"duration":4.066261,"end_time":"2025-06-25T09:11:57.206636","exception":false,"start_time":"2025-06-25T09:11:53.140375","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport ast\n\nVISUALIZE_SAMPLE = True  # ✅ 控制是否执行这段单样本评估\n\npred_csv = \"/kaggle/working/cell_level_predictions_conf_05_top4_train.csv\"\ntrue_csv = \"/kaggle/working/bbox_with_label_vector.csv\"\n\nif VISUALIZE_SAMPLE:\n    pred_df = pd.read_csv(pred_csv)\n    true_df = pd.read_csv(true_csv)\n\n    # Step 1: 随机选择一个 cell\n    sample_row = pred_df.sample(1).iloc[0]\n    image_id = sample_row[\"image_id\"]\n    cell_id = sample_row[\"cell_id\"]\n\n    print(f\"📸 Image ID: {image_id}\")\n    print(f\"🧫 Cell ID: {cell_id}\")\n\n    # Step 2: 获取预测类别\n    cell_preds = pred_df[(pred_df[\"image_id\"] == image_id) & (pred_df[\"cell_id\"] == cell_id)]\n    pred_classes = sorted(cell_preds[\"class_id\"].unique())\n    print(f\"✅ Predicted class IDs: {pred_classes}\")\n\n    # Step 3: 获取真实标签向量\n    true_row = true_df[(true_df[\"image_id\"] == image_id) & (true_df[\"cell_id\"] == cell_id)]\n    if not true_row.empty:\n        label_vector = ast.literal_eval(true_row.iloc[0][\"label_vector\"])\n        true_classes = [i for i, v in enumerate(label_vector) if v == 1]\n        print(f\"🎯 True class IDs: {true_classes}\")\n    else:\n        label_vector = [0] * 19\n        true_classes = []\n        print(\"⚠️ 无真实标签\")\n\n    # Step 4: 从 dataset 中获取样本\n    assert 'dataset' in globals(), \"❌ dataset 未定义，请先构建 WeakSupervisionDataset\"\n    idx = dataset.df[(dataset.df[\"image_id\"] == image_id) & (dataset.df[\"cell_id\"] == cell_id)].index\n    if len(idx) == 0:\n        print(\"❌ 在 dataset 中找不到该样本\")\n    else:\n        img, mask, label_tensor, _, _ = dataset[idx[0]]\n        img_np = img[:3].permute(1, 2, 0).cpu().numpy()\n        img_np = (img_np * 0.5 + 0.5).clip(0, 1)\n\n        # Step 5: 模型推理\n        img_tensor = img.unsqueeze(0)[:, :3].to(DEVICE)\n        with torch.no_grad():\n            pred_logits = model(img_tensor)\n            pred_probs = torch.sigmoid(pred_logits[0])\n\n        # Step 6: 可视化\n        num_classes = len(pred_classes)\n        if num_classes == 0:\n            print(\"❌ 无预测类别，跳过可视化\")\n        else:\n            plt.figure(figsize=(4 * (num_classes + 1), 4))\n            plt.subplot(1, num_classes + 1, 1)\n            plt.imshow(img_np)\n            plt.title(\"Cropped Cell\")\n            plt.axis(\"off\")\n\n            for idx_plot, class_id in enumerate(pred_classes):\n                prob_val = pred_probs[class_id].item()\n                plt.subplot(1, num_classes + 1, idx_plot + 2)\n                plt.imshow(img_np)\n                plt.title(f\"Class {class_id}\\nConf: {prob_val:.2f}\")\n                plt.axis(\"off\")\n\n            plt.tight_layout()\n            plt.show()\n\n        # Step 7: 评估指标\n        true_vec = label_tensor.cpu().numpy().astype(int)\n        pred_vec = np.zeros(19, dtype=int)\n        pred_vec[pred_classes] = 1\n\n        intersection = np.logical_and(pred_vec, true_vec).sum()\n        union = np.logical_or(pred_vec, true_vec).sum()\n        jaccard = intersection / union if union > 0 else 0.0\n        cosine = np.dot(pred_vec, true_vec) / (np.linalg.norm(pred_vec) * np.linalg.norm(true_vec) + 1e-6)\n        accuracy = (pred_vec == true_vec).sum() / 19\n\n        print(f\"\\n📊 Jaccard (IoU): {jaccard:.4f}\")\n        print(f\"📊 Cosine similarity: {cosine:.4f}\")\n        print(f\"📊 Accuracy (full 19-class vector): {accuracy:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-24T03:10:27.1219Z","iopub.execute_input":"2025-07-24T03:10:27.122191Z","iopub.status.idle":"2025-07-24T03:10:28.14542Z","shell.execute_reply.started":"2025-07-24T03:10:27.122172Z","shell.execute_reply":"2025-07-24T03:10:28.144533Z"}},"outputs":[],"execution_count":null}]}