{"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"}],"dockerImageVersionId":31089,"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-22T10:04:56.546355Z","iopub.execute_input":"2025-07-22T10:04:56.547107Z","iopub.status.idle":"2025-07-22T10:04:56.551655Z","shell.execute_reply.started":"2025-07-22T10:04:56.547083Z","shell.execute_reply":"2025-07-22T10:04:56.550605Z"}},"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 = 5\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-22T10:04:56.552886Z","iopub.execute_input":"2025-07-22T10:04:56.553356Z","iopub.status.idle":"2025-07-22T10:04:56.58207Z","shell.execute_reply.started":"2025-07-22T10:04:56.553328Z","shell.execute_reply":"2025-07-22T10:04:56.581478Z"}},"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-22T10:04:56.582679Z","iopub.execute_input":"2025-07-22T10:04:56.582838Z","iopub.status.idle":"2025-07-22T10:04:57.090254Z","shell.execute_reply.started":"2025-07-22T10:04:56.582826Z","shell.execute_reply":"2025-07-22T10:04:57.089567Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 读取train and bbox csv","metadata":{}},{"cell_type":"code","source":"import pandas as pd\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# 读取两个 CSV\nbbox_df = pd.read_csv(bbox_path)\ntrain_df = pd.read_csv(train_path)\n\n# 显示变量名称和前几行\nprint(\"========== bbox_filtered.csv ==========\")\nprint(\"作用：提供每个 image_id 下每个 cell 的 bounding box、cell_id、cell 大小等信息\")\nprint(\"列名如下：\")\nprint(bbox_df.columns.tolist())\nprint(\"\\n示例数据（前5行）：\")\nprint(bbox_df.head())\n\nprint(\"\\n========== train.csv ==========\")\nprint(\"作用：提供每张图像（image_id）对应的整体标签（图像级）\")\nprint(\"列名如下：\")\nprint(train_df.columns.tolist())\nprint(\"\\n示例数据（前5行）：\")\nprint(train_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:04:57.092317Z","iopub.execute_input":"2025-07-22T10:04:57.092527Z","iopub.status.idle":"2025-07-22T10:04:57.735454Z","shell.execute_reply.started":"2025-07-22T10:04:57.09251Z","shell.execute_reply":"2025-07-22T10:04:57.734737Z"}},"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-22T10:04:57.736249Z","iopub.execute_input":"2025-07-22T10:04:57.736505Z","iopub.status.idle":"2025-07-22T10:05:01.440016Z","shell.execute_reply.started":"2025-07-22T10:04:57.736485Z","shell.execute_reply":"2025-07-22T10:05:01.439267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bbox_df = pd.read_csv(\"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/train_bbox_filtered.csv\")\ntrain_df = pd.read_csv(\"/kaggle/input/hpa-single-cell-image-classification/train.csv\")\n\n# 获取 ID 集合\nbbox_ids = set(bbox_df[\"image_id\"].unique())\ntrain_ids = set(train_df[\"ID\"].unique())\n\n# 计算匹配度\nmatched_ids = bbox_ids & train_ids\n\nprint(f\"总共有 {len(bbox_ids)} 个 bbox 图像ID\")\nprint(f\"总共有 {len(train_ids)} 个 train 图像ID\")\nprint(f\"交集图像ID数: {len(matched_ids)}，匹配比例: {len(matched_ids)/len(bbox_ids)*100:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:05:01.440802Z","iopub.execute_input":"2025-07-22T10:05:01.441124Z","iopub.status.idle":"2025-07-22T10:05:02.077205Z","shell.execute_reply.started":"2025-07-22T10:05:01.441103Z","shell.execute_reply":"2025-07-22T10:05:02.076492Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 看一下npz文件里都有什么并创建一个目录¶","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\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)\n\n# 获取唯一字段 arr_0\ncontent = data[\"arr_0\"]\n\n# 打印结构信息\nprint(\"字段类型：\", type(content))\nprint(\"元素数量（如有）：\", len(content))\n\n# 打印前几个内容查看结构\nif isinstance(content, dict):\n    print(\"dict keys 示例：\", list(content.keys())[:3])\nelif isinstance(content, list):\n    print(\"list 前3项：\", content[:3])\nelif isinstance(content, np.ndarray):\n    print(\"ndarray shape：\", content.shape)\n    print(\"前1个元素类型：\", type(content[0]))\n    print(\"前1个元素内容（简略）：\", content[0])\nelse:\n    print(\"未知类型：\", content)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:05:02.078006Z","iopub.execute_input":"2025-07-22T10:05:02.078717Z","iopub.status.idle":"2025-07-22T10:05:02.1067Z","shell.execute_reply.started":"2025-07-22T10:05:02.078688Z","shell.execute_reply":"2025-07-22T10:05:02.105989Z"}},"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):\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\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, (target_w, target_h))\n        y = cv2.resize(y, (target_w, target_h))\n        b = cv2.resize(b, (target_w, 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, (target_w, target_h), interpolation=cv2.INTER_NEAREST)\n\n        # === Step 4: ⚠️ 同步缩放 bbox 坐标 ===\n        scale_x = target_w / original_w\n        scale_y = 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: 构造正方 bbox 并加 padding ===\n        H, W = image.shape[:2]\n        bbox_w = x1 - x0\n        bbox_h = y1 - y0\n        side = max(bbox_w, bbox_h)\n        pad = 10\n\n        cx = (x0 + x1) // 2\n        cy = (y0 + y1) // 2\n\n        new_x0 = max(0, cx - side // 2 - pad)\n        new_x1 = min(W, cx + side // 2 + pad)\n        new_y0 = max(0, cy - side // 2 - pad)\n        new_y1 = min(H, cy + side // 2 + pad)\n\n        # === Step 6: 异常处理 ===\n        if new_x1 <= new_x0 or new_y1 <= new_y0:\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            return crop_tensor, label_tensor, image_id, cell_id\n\n        # === Step 7: crop 图像并 resize ===\n        crop_img = image[new_y0:new_y1, new_x0:new_x1, :]\n        crop_img = cv2.resize(crop_img, (self.crop_size, self.crop_size))\n        crop_img = Image.fromarray(crop_img.astype(np.uint8))  # 转为 PIL Image\n\n        # === Step 8: 图像 transform / 归一化 ===\n        if self.transform:\n            crop_tensor = self.transform(crop_img)\n        else:\n            crop_np = np.array(crop_img).astype(np.float32) / 255.0\n            crop_np = np.transpose(crop_np, (2, 0, 1))\n            crop_tensor = torch.tensor(crop_np).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, label_tensor, image_id, cell_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:25:04.606378Z","iopub.execute_input":"2025-07-22T10:25:04.607048Z","iopub.status.idle":"2025-07-22T10:25:04.61904Z","shell.execute_reply.started":"2025-07-22T10:25:04.607021Z","shell.execute_reply":"2025-07-22T10:25:04.618056Z"}},"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\nfrom torchvision import transforms\nimport pandas as pd\n\n# 加载 CSV（可选）\ndf_valid = pd.read_csv('/kaggle/working/bbox_with_label_vector.csv')\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntransform = A.Compose([\n    A.Resize(256, 256),  # 或 512×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=\"/kaggle/working/bbox_with_label_vector.csv\",\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\n# ✅ 创建 dataloader\ntrain_loader = DataLoader(dataset, batch_size=32, shuffle=True)\n\n# ✅ 检查第一张样本（变量名顺序更新）\nimage_tensor, mask_tensor, label_tensor, image_id, cell_id = dataset[0]\n\nprint(f\"图像大小: {image_tensor.shape}\")\nprint(f\"掩膜大小: {mask_tensor.shape}\")\nprint(f\"标签向量: {label_tensor}\")\nprint(f\"图像 ID: {image_id}, Cell ID: {cell_id}\")\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-22T10:30:53.940343Z","iopub.execute_input":"2025-07-22T10:30:53.940768Z","iopub.status.idle":"2025-07-22T10:30:55.52773Z","shell.execute_reply.started":"2025-07-22T10:30:53.940743Z","shell.execute_reply":"2025-07-22T10:30:55.526957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\ndf_valid = pd.read_csv('/kaggle/working/bbox_with_label_vector.csv')\n\n# 解析并求和\nlabel_matrix = df_valid['label_vector'].apply(lambda x: np.array(eval(x)))  # CSV中是列表格式字符串\nlabel_sum = np.sum(np.stack(label_matrix.values), axis=0)\n\nprint(\"每个类别的样本数:\", label_sum.astype(int))\n","metadata":{"papermill":{"duration":0.419357,"end_time":"2025-06-25T08:04:05.910361","exception":false,"start_time":"2025-06-25T08:04:05.491004","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:07.139423Z","iopub.execute_input":"2025-07-22T10:31:07.140188Z","iopub.status.idle":"2025-07-22T10:31:09.697685Z","shell.execute_reply.started":"2025-07-22T10:31:07.140161Z","shell.execute_reply":"2025-07-22T10:31:09.696698Z"}},"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\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# 使用 Albumentations 做图像增强 + resize\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 实例（注意 crop_size 改为 target_size）\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  # 原图 resize 尺寸，mask/image 都会先 resize 到这个大小再 crop\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:12.784972Z","iopub.execute_input":"2025-07-22T10:31:12.785607Z","iopub.status.idle":"2025-07-22T10:31:13.430697Z","shell.execute_reply.started":"2025-07-22T10:31:12.785581Z","shell.execute_reply":"2025-07-22T10:31:13.430056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\n\ndef visualize_random_crops(dataset, num_samples=12, rows=3, 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        image_tensor, mask_tensor, label_tensor, image_id, cell_id = dataset[idx]\n        \n        # 图像还原为 numpy\n        img_np = image_tensor.permute(1, 2, 0).cpu().numpy()  # [H, W, C]\n        img_np = (img_np * 0.5 + 0.5).clip(0, 1)  # 去 Normalize\n\n        # mask 转 numpy（确保是 [H,W]）\n        mask_np = mask_tensor.squeeze().cpu().numpy()\n\n        # 显示图像 + 掩膜 overlay\n        ax.imshow(img_np)\n        ax.imshow(mask_np, cmap='Reds', alpha=0.4)\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# ✅ 调用方法（假设你的 dataset 叫做 train_dataset）\nvisualize_random_crops(train_dataset)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:17.039197Z","iopub.execute_input":"2025-07-22T10:31:17.039718Z","iopub.status.idle":"2025-07-22T10:31:21.339227Z","shell.execute_reply.started":"2025-07-22T10:31:17.039694Z","shell.execute_reply":"2025-07-22T10:31:21.338301Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"在这里我们使用了旋转，让图像可以得到更多的训练","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport cv2\nimport numpy as np\n\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        # 获取样本数据\n        image_tensor, mask_tensor, label_vector, image_id, cell_id = dataset[idx]\n\n        # 图像还原\n        crop_img = image_tensor.permute(1, 2, 0).numpy()  # [H,W,C]\n        crop_img = (crop_img * 0.5 + 0.5).clip(0, 1)\n\n        # 掩膜还原\n        mask_np = mask_tensor.squeeze().numpy()\n\n        # === Step 1: 加载原图并 resize 到 2048 ===\n        r = cv2.imread(f\"/kaggle/input/hpa-single-cell-image-classification/train/{image_id}_red.png\", cv2.IMREAD_GRAYSCALE)\n        y = cv2.imread(f\"/kaggle/input/hpa-single-cell-image-classification/train/{image_id}_yellow.png\", cv2.IMREAD_GRAYSCALE)\n        b = cv2.imread(f\"/kaggle/input/hpa-single-cell-image-classification/train/{image_id}_blue.png\", cv2.IMREAD_GRAYSCALE)\n        rgb = np.stack([r, y, b], axis=-1)\n        rgb = cv2.resize(rgb, (2048, 2048))\n        rgb = rgb.astype(np.float32) / 255.0\n\n        # === Step 2: 加载 mask 并提取 cell 的区域 ===\n        mask_path = f\"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/mask/train/cell/{image_id}.npz\"\n        full_mask = np.load(mask_path, 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\n        # === Step 3: overlay mask 到原图 ===\n        overlay = rgb.copy()\n        overlay[cell_mask == 1] = [1.0, 0.0, 0.0]  # 红色区域高亮\n\n        # === Step 4: 显示两张图 ===\n        plt.figure(figsize=(10, 4))\n\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\n        plt.tight_layout()\n        plt.show()\n\nvisualize_full_and_crop(train_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:28.571Z","iopub.execute_input":"2025-07-22T10:31:28.571774Z","iopub.status.idle":"2025-07-22T10:31:35.162989Z","shell.execute_reply.started":"2025-07-22T10:31:28.571746Z","shell.execute_reply":"2025-07-22T10:31:35.162249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# los_weight","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 pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\n# 读取图像级标签数据\nmeta_path = '/kaggle/input/hpa-single-cell-image-classification/train.csv'\ndf_meta = pd.read_csv(meta_path)\n\n# 拆分标签列为 list（每行是一个 label list）\ndf_meta['Label_list'] = df_meta['Label'].apply(lambda x: list(map(int, x.split('|'))))\n\n# 初始化每类计数数组\nNUM_CLASSES = 19\nlabel_counts = np.zeros(NUM_CLASSES, dtype=int)\n\n# 统计每个类别在 image-level 上的出现次数\nfor labels in df_meta['Label_list']:\n    for cls in labels:\n        label_counts[cls] += 1\n\nprint(\"每个类别的 image-level 出现次数:\", label_counts)\n# 原始频率计算\nclass_freq = label_counts / label_counts.sum()\n\n# 对数缩放，避免 weight 差距过大\npos_weight = np.log((1 - class_freq) / class_freq + 1e-8)\n\n# 对于 class 18 和 class 0、16 手动调整\npos_weight[18] = 0.0\npos_weight[0] = min(pos_weight[0], 3.5)   # 降低一点惩罚\npos_weight[16] = min(pos_weight[16], 3)\n\n# 打印每个类别的 pos_weight\nprint(\"\\n每个类别的 BCEWithLogitsLoss 权重（pos_weight）如下：\")\nfor i in range(NUM_CLASSES):\n    print(f\"class {i:2d}: pos_weight = {pos_weight[i]:.4f}\")","metadata":{"papermill":{"duration":0.081858,"end_time":"2025-06-25T08:04:06.0707","exception":false,"start_time":"2025-06-25T08:04:05.988842","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:40.652533Z","iopub.execute_input":"2025-07-22T10:31:40.652869Z","iopub.status.idle":"2025-07-22T10:31:40.714249Z","shell.execute_reply.started":"2025-07-22T10:31:40.652845Z","shell.execute_reply":"2025-07-22T10:31:40.713592Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# define model and loss function and optimizer","metadata":{"papermill":{"duration":0.015811,"end_time":"2025-06-25T08:04:06.102771","exception":false,"start_time":"2025-06-25T08:04:06.08696","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from torchvision import models\nimport torch\nimport torch.nn as nn\n\n# 构建模型\nfrom torchvision.models import ResNet18_Weights\nmodel = models.resnet18(weights=None)  # 等价于不加载预训练参数\nmodel.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)\nmodel = model.to(DEVICE)\n\n# 设置 class 18 不参与训练\npos_weight[18] = 0.0\n\n# 定义加权 BCE 损失\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(model.parameters(), lr=LR)","metadata":{"papermill":{"duration":4.489641,"end_time":"2025-06-25T08:04:10.608253","exception":false,"start_time":"2025-06-25T08:04:06.118612","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:31:44.921246Z","iopub.execute_input":"2025-07-22T10:31:44.92202Z","iopub.status.idle":"2025-07-22T10:31:45.091934Z","shell.execute_reply.started":"2025-07-22T10:31:44.921983Z","shell.execute_reply":"2025-07-22T10:31:45.091104Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# loop train","metadata":{"papermill":{"duration":0.016038,"end_time":"2025-06-25T08:04:10.640961","exception":false,"start_time":"2025-06-25T08:04:10.624923","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch\nimport os\nimport numpy as np\n\n# 设置保存模型的路径\nbest_val_loss = float('inf')\nbest_model_path = \"/kaggle/working/best_model_train.pth\"\n\n# 训练轮数（你可以自定义）\nEPOCHS = 5\n\nfor epoch in range(EPOCHS):\n    model.train()\n    total_loss = 0.0\n\n    train_loader_tqdm = tqdm(train_loader, desc=f\"[Train] Epoch {epoch+1}\")\n    step = 0  # ✅ 初始化 step\n\n    for images, _, labels, image_ids, cell_ids in train_loader_tqdm:\n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)\n\n        # 前向传播\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\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 > 4:  # ✅ 最多只训练 16 个 batch（32×16=512 个样本）\n            break\n\n    avg_train_loss = total_loss / step  # 注意除以 step 而不是 len(train_loader)\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\n        avg_val_loss = val_loss / val_step\n        print(f\"Validation loss: {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\"✅ 新最佳模型已保存到: {best_model_path}（val_loss: {best_val_loss:.4f}）\")\n\n    else:\n        # 没有验证集，基于训练 loss 保存\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\"✅ 新最佳模型（基于训练损失）已保存到: {best_model_path}\")","metadata":{"papermill":{"duration":3374.241999,"end_time":"2025-06-25T09:00:24.899113","exception":false,"start_time":"2025-06-25T08:04:10.657114","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:39:31.837541Z","iopub.execute_input":"2025-07-22T10:39:31.837858Z","iopub.status.idle":"2025-07-22T10:41:55.270162Z","shell.execute_reply.started":"2025-07-22T10:39:31.837834Z","shell.execute_reply":"2025-07-22T10:41:55.269348Z"}},"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\n\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model_train.pth\", map_location=DEVICE))\nmodel.eval()\n\nresults = []\nthreshold = 0.2\nNUM_CLASSES = 19\nmax_steps = 5  # 只处理前 30 个 batch\n\nwith torch.no_grad():\n    for step, (images, _, labels, image_ids, cell_ids) in enumerate(tqdm(train_loader, desc=\"Collecting predictions\")):\n        if step >= max_steps:\n            break\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\ndf_result = pd.DataFrame(results)\ndf_result.to_csv(\"/kaggle/working/cell_level_predictions_conf_07_train.csv\", index=False)\nprint(\"✅ 保存成功，共有预测类别数:\", 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-22T10:41:55.271538Z","iopub.execute_input":"2025-07-22T10:41:55.271797Z","iopub.status.idle":"2025-07-22T10:42:28.639928Z","shell.execute_reply.started":"2025-07-22T10:41:55.271778Z","shell.execute_reply":"2025-07-22T10:42:28.639091Z"}},"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_07_train.csv')\n\n# 显示前 15 行\nprint(df.head(15))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T10:42:28.64084Z","iopub.execute_input":"2025-07-22T10:42:28.641045Z","iopub.status.idle":"2025-07-22T10:42:28.651809Z","shell.execute_reply.started":"2025-07-22T10:42:28.64103Z","shell.execute_reply":"2025-07-22T10:42:28.651117Z"}},"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\n\ndf = pd.read_csv('/kaggle/working/cell_level_predictions_conf_07_train.csv')\n\n# 对 image_id 分组，处理每个 group\ndef keep_top4_classes(group):\n    # 获取该 image_id 的唯一 class_id 个数\n    if len(group['class_id'].unique()) > 4:\n        # 按 confidence 降序排序，保留 top 4 class\n        return group.sort_values('confidence', ascending=False).drop_duplicates('class_id').head(4)\n    else:\n        return group\n\n# 按 image_id 分组应用筛选逻辑\ndf_top4 = df.groupby('image_id', group_keys=False).apply(keep_top4_classes).reset_index(drop=True)\n\n# 保存结果\ndf_top4.to_csv('/kaggle/working/cell_level_predictions_conf_07_top4_train.csv', index=False)\nprint(\"已完成按 image_id 保留 top 4 class 的筛选，结果已保存至:\")\nprint(\"/kaggle/working/cell_level_predictions_conf_07_top4_train.csv\")","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-22T10:42:28.652653Z","iopub.execute_input":"2025-07-22T10:42:28.653172Z","iopub.status.idle":"2025-07-22T10:42:28.762469Z","shell.execute_reply.started":"2025-07-22T10:42:28.653146Z","shell.execute_reply":"2025-07-22T10:42:28.761715Z"}},"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_07_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-22T10:42:28.764267Z","iopub.execute_input":"2025-07-22T10:42:28.765102Z","iopub.status.idle":"2025-07-22T10:42:35.584595Z","shell.execute_reply.started":"2025-07-22T10:42:28.765068Z","shell.execute_reply":"2025-07-22T10:42:35.58378Z"}},"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_07_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-22T10:42:35.585314Z","iopub.execute_input":"2025-07-22T10:42:35.585579Z","iopub.status.idle":"2025-07-22T10:42:35.594454Z","shell.execute_reply.started":"2025-07-22T10:42:35.585548Z","shell.execute_reply":"2025-07-22T10:42:35.593602Z"}},"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-22T10:42:35.595174Z","iopub.execute_input":"2025-07-22T10:42:35.595368Z","iopub.status.idle":"2025-07-22T10:42:36.238877Z","shell.execute_reply.started":"2025-07-22T10:42:35.595348Z","shell.execute_reply":"2025-07-22T10:42:36.238127Z"}},"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_07_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-22T10:42:36.239489Z","iopub.execute_input":"2025-07-22T10:42:36.239706Z","iopub.status.idle":"2025-07-22T10:42:36.397251Z","shell.execute_reply.started":"2025-07-22T10:42:36.239689Z","shell.execute_reply":"2025-07-22T10:42:36.396458Z"}},"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_07_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-22T10:42:36.398176Z","iopub.execute_input":"2025-07-22T10:42:36.398474Z","iopub.status.idle":"2025-07-22T10:42:36.40927Z","shell.execute_reply.started":"2025-07-22T10:42:36.398448Z","shell.execute_reply":"2025-07-22T10:42:36.408678Z"}},"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\n\n# ========= 参数和文件路径 =========\npred_csv = \"/kaggle/working/cell_level_predictions_conf_07_top4_train.csv\"\ntrue_csv = \"/kaggle/working/bbox_with_label_vector.csv\"\n\n# ========= 读取数据 =========\npred_df = pd.read_csv(pred_csv)\ntrue_df = pd.read_csv(true_csv)\n\n# ========= Step 1: 随机选择一组 (image_id, cell_id) =========\nsample_row = pred_df.sample(1).iloc[0]\nimage_id = sample_row[\"image_id\"]\ncell_id = sample_row[\"cell_id\"]\n\n# ========= Step 2: 获取预测类别 =========\ncell_preds = pred_df[(pred_df[\"image_id\"] == image_id) & (pred_df[\"cell_id\"] == cell_id)]\npred_classes = sorted(cell_preds[\"class_id\"].unique())\n\nprint(f\"📸 Image ID: {image_id}\")\nprint(f\"🧫 Cell ID: {cell_id}\")\nprint(f\"✅ Predicted class IDs: {pred_classes}\")\n\n# ========= Step 3: 获取真实标签向量 =========\ntrue_row = true_df[(true_df[\"image_id\"] == image_id) & (true_df[\"cell_id\"] == cell_id)]\nif not true_row.empty:\n    label_vector = 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}\")\nelse:\n    true_classes = []\n    label_vector = [0] * 19\n    print(\"⚠️ 无真实标签\")\n\n# ========= Step 4: 从 dataset 中获取该样本 =========\nidx = dataset.df[(dataset.df[\"image_id\"] == image_id) & (dataset.df[\"cell_id\"] == cell_id)].index\nif len(idx) == 0:\n    print(\"❌ 无法在 dataset 中找到样本\")\nelse:\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\n        # 原图\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-22T10:43:41.233294Z","iopub.execute_input":"2025-07-22T10:43:41.233903Z","iopub.status.idle":"2025-07-22T10:43:42.479498Z","shell.execute_reply.started":"2025-07-22T10:43:41.233878Z","shell.execute_reply":"2025-07-22T10:43:42.478725Z"}},"outputs":[],"execution_count":null}]}