{"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":12067245,"sourceType":"datasetVersion","datasetId":7595523},{"sourceId":12124945,"sourceType":"datasetVersion","datasetId":7634750},{"sourceId":244656261,"sourceType":"kernelVersion"},{"sourceId":251505556,"sourceType":"kernelVersion"}],"dockerImageVersionId":30056,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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.838023,"end_time":"2021-01-31T14:33:44.653299","exception":false,"start_time":"2021-01-31T14:33:42.815276","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:30.387021Z","iopub.execute_input":"2025-07-21T11:50:30.387353Z","iopub.status.idle":"2025-07-21T11:50:30.773866Z","shell.execute_reply.started":"2025-07-21T11:50:30.387281Z","shell.execute_reply":"2025-07-21T11:50:30.772981Z"}},"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)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:30.775725Z","iopub.execute_input":"2025-07-21T11:50:30.775962Z","iopub.status.idle":"2025-07-21T11:50:30.809182Z","shell.execute_reply.started":"2025-07-21T11:50:30.775939Z","shell.execute_reply":"2025-07-21T11:50:30.808496Z"}},"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()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:30.810406Z","iopub.execute_input":"2025-07-21T11:50:30.810655Z","iopub.status.idle":"2025-07-21T11:50:31.296541Z","shell.execute_reply.started":"2025-07-21T11:50:30.810628Z","shell.execute_reply":"2025-07-21T11:50:31.295721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取 train.csv\ndf = pd.read_csv(\"/kaggle/input/hpa-single-cell-image-classification/train.csv\")\n\n# 查找该图像的标签\nimage_id = \"0042017c-bba4-11e8-b2b9-ac1f6b6435d0\"\nrecord = df[df[\"ID\"] == image_id]\n\n# 显示结果\nprint(record)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:31.298283Z","iopub.execute_input":"2025-07-21T11:50:31.298709Z","iopub.status.idle":"2025-07-21T11:50:31.393064Z","shell.execute_reply.started":"2025-07-21T11:50:31.298666Z","shell.execute_reply":"2025-07-21T11:50:31.392187Z"}},"outputs":[],"execution_count":null},{"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())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:31.3953Z","iopub.execute_input":"2025-07-21T11:50:31.395539Z","iopub.status.idle":"2025-07-21T11:50:32.00722Z","shell.execute_reply.started":"2025-07-21T11:50:31.395516Z","shell.execute_reply":"2025-07-21T11:50:32.006456Z"}},"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# 构建：image_id → list of label indices 的映射\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)\nimageid_to_labels = dict(zip(train_df['ID'], train_df['label_list']))\n\n# 准备输出 DataFrame\nrecords = []\n\nfor idx, row in bbox_df.iterrows():\n    image_id = row['image_id']\n    cell_id = row['cell_id']\n    \n    # 如果图像ID在标签文件中，就生成multi-hot label\n    label_indices = imageid_to_labels.get(image_id, [])\n    label_vector = [0] * 19\n    for i in label_indices:\n        if 0 <= i < 19:\n            label_vector[i] = 1\n            \n    records.append({\n        'image_id': image_id,\n        'cell_id': cell_id,\n        'label_vector': label_vector\n    })\n\n# 保存结果\noutput_df = pd.DataFrame(records)\noutput_df.to_csv(\"/kaggle/working/weak_supervision_labels.csv\", index=False)\n\nprint(\"保存成功！共 %d 个 cell-level 弱标签样本。\" % len(output_df))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:32.008788Z","iopub.execute_input":"2025-07-21T11:50:32.009014Z","iopub.status.idle":"2025-07-21T11:50:37.893091Z","shell.execute_reply.started":"2025-07-21T11:50:32.008991Z","shell.execute_reply":"2025-07-21T11:50:37.892315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 再次读取\nbbox_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}%\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:37.894303Z","iopub.execute_input":"2025-07-21T11:50:37.894624Z","iopub.status.idle":"2025-07-21T11:50:38.501985Z","shell.execute_reply.started":"2025-07-21T11:50:37.894595Z","shell.execute_reply":"2025-07-21T11:50:38.501164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf_weak = pd.read_csv(\"/kaggle/working/weak_supervision_labels.csv\")\n\n# 打印变量名（字段名）与前几行\nprint(\" CSV 中的字段名（变量名）：\")\nprint(df_weak.columns.tolist())\n\nprint(\"\\n 前几行数据示例：\")\nprint(df_weak.head(3))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:38.503057Z","iopub.execute_input":"2025-07-21T11:50:38.503281Z","iopub.status.idle":"2025-07-21T11:50:38.563703Z","shell.execute_reply.started":"2025-07-21T11:50:38.503258Z","shell.execute_reply":"2025-07-21T11:50:38.562859Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 设置参数","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport os\nimport torch.nn.functional as F\n\nclass WeakSupervisionCellDataset(torch.utils.data.Dataset):\n    def __init__(self, df, image_dir, mask_dir, transform=None, target_size=2048, num_classes=19):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir  # 存放 RGBY 四张图像的目录\n        self.mask_dir = mask_dir    # 存放每个图像的 .npz 掩膜文件\n        self.transform = transform\n        self.target_size = target_size\n        self.num_classes = num_classes\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        label_vector = torch.tensor(eval(row[\"label_vector\"]), dtype=torch.float32)\n\n        # === Step 1: 读取 .npz 掩膜并 resize 到 2048x2048 ===\n        mask_path = os.path.join(self.mask_dir, f\"{image_id}.npz\")\n        mask_np = np.load(mask_path, allow_pickle=True)[\"arr_0\"]  # 原始 shape\n        original_h, original_w = mask_np.shape\n        scale_x = self.target_size / original_w\n        scale_y = self.target_size / original_h\n        mask_resized = cv2.resize(mask_np, (self.target_size, self.target_size), interpolation=cv2.INTER_NEAREST)\n\n        # === Step 2: 获取 cell bbox，计算新坐标 ===\n        x0, y0 = int(row[\"x0\"]), int(row[\"y0\"])\n        w, h = int(row[\"w_cell\"]), int(row[\"h_cell\"])\n        x1, y1 = x0 + w, y0 + h\n        x0_new = int(x0 * scale_x)\n        y0_new = int(y0 * scale_y)\n        x1_new = int(x1 * scale_x)\n        y1_new = int(y1 * scale_y)\n\n        # === Step 3: 裁剪 mask ===\n        cell_mask = mask_resized[y0_new:y1_new, x0_new:x1_new]  # shape: (h, w)\n\n        # === Step 4: 读取 RGBY 图像 & 裁剪 ===\n        def load_channel(channel):\n            path = os.path.join(self.image_dir, f\"{image_id}_{channel}.png\")\n            img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                raise FileNotFoundError(f\"{channel} 图像不存在: {path}\")\n            img_resized = cv2.resize(img, (self.target_size, self.target_size))\n            return img_resized[y0_new:y1_new, x0_new:x1_new]\n\n        r = load_channel(\"red\")\n        y = load_channel(\"yellow\")\n        b = load_channel(\"blue\")\n        g = y  # ✅ 将 yellow 替代 green 通道\n\n        image = np.stack([r, g, b], axis=-1)  # shape: (h, w, 3)\n\n        # === Step 5: 数据增强（如果有）===\n        if self.transform:\n            transformed = self.transform(image=image, mask=cell_mask)\n            image = transformed[\"image\"]\n            cell_mask = transformed[\"mask\"]\n\n        # === Step 6: 转 tensor ===\n        image = torch.tensor(image).permute(2, 0, 1).float() / 255.0  # (3, H, W)\n        cell_mask = torch.tensor(cell_mask, dtype=torch.long)        # (H, W)\n        label_vector = torch.tensor(eval(row[\"label_vector\"]), dtype=torch.float32)\n\n        return image, cell_mask, label_vector, image_id, cell_id\n\n    def __len__(self):\n        return len(self.df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:38.564843Z","iopub.execute_input":"2025-07-21T11:50:38.565087Z","iopub.status.idle":"2025-07-21T11:50:38.581317Z","shell.execute_reply.started":"2025-07-21T11:50:38.565062Z","shell.execute_reply":"2025-07-21T11:50:38.580636Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据集构造","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# ======================\n# Step 1: 准备 Dataset 类\n# ======================\nclass CellSegmentationDataset(Dataset):\n    def __init__(self, label_csv, bbox_csv, image_dir, mask_dir, transform=None):\n        self.df_label = pd.read_csv(label_csv)\n        self.df_bbox = pd.read_csv(bbox_csv)\n        self.df = pd.merge(self.df_label, self.df_bbox, on=[\"image_id\", \"cell_id\"])\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.transform = transform\n        self.num_classes = 19\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 = int(row[\"cell_id\"])\n        x0, y0, w, h = int(row[\"x0\"]), int(row[\"y0\"]), int(row[\"w_cell\"]), int(row[\"h_cell\"])\n        x1, y1 = x0 + w, y0 + h\n\n        # 读取三个通道图像：R, Y, B\n        def read_gray(path): return cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        r = read_gray(os.path.join(self.image_dir, f\"{image_id}_red.png\"))\n        y = read_gray(os.path.join(self.image_dir, f\"{image_id}_yellow.png\"))  # 用 yellow 替代 green\n        b = read_gray(os.path.join(self.image_dir, f\"{image_id}_blue.png\"))\n        full_image = np.stack([r, y, b], axis=-1)\n\n        # mask 解码\n        mask_path = os.path.join(self.mask_dir, f\"{image_id}.npz\")\n        full_mask = np.load(mask_path, allow_pickle=True)[\"arr_0\"]\n\n        # resize 到 2048×2048\n        TARGET_SIZE = 2048\n        full_image = cv2.resize(full_image, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_LINEAR)\n        full_mask = cv2.resize(full_mask, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_NEAREST)\n\n        # 计算坐标比例并截取\n        scale_x, scale_y = TARGET_SIZE / r.shape[1], TARGET_SIZE / r.shape[0]\n        x0, y0 = int(x0 * scale_x), int(y0 * scale_y)\n        x1, y1 = int(x1 * scale_x), int(y1 * scale_y)\n\n        crop_img = full_image[y0:y1, x0:x1]\n        crop_mask = (full_mask[y0:y1, x0:x1] == cell_id).astype(np.uint8)  # 仅保留该 cell 的 mask\n\n        # 图像增强\n        if self.transform:\n            augmented = self.transform(image=crop_img, mask=crop_mask)\n            crop_img = augmented[\"image\"]\n            crop_mask = augmented[\"mask\"]\n\n        # 将 mask 转为 tensor，并添加 channel 维度\n        crop_mask = torch.tensor(crop_mask, dtype=torch.long).unsqueeze(0)  # shape: (1, H, W)\n        # label vector\n        label_vector = np.array(eval(row[\"label_vector\"]), dtype=np.float32)\n\n        return crop_img, crop_mask, label_vector, image_id, cell_id\n\n# 暂不训练，只返回 dataset size 确认\nlabel_csv = \"/kaggle/working/weak_supervision_labels.csv\"\nbbox_csv = \"/kaggle/input/hpa2021-p-all-1stx8-2ndx16-imgx16/train_bbox_filtered.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.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),\n    ToTensorV2()\n])\n\ndataset = CellSegmentationDataset(label_csv, bbox_csv, image_dir, mask_dir, transform=transform)\n\nprint(dataset.df.head(5))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:38.582869Z","iopub.execute_input":"2025-07-21T11:50:38.583271Z","iopub.status.idle":"2025-07-21T11:50:39.793592Z","shell.execute_reply.started":"2025-07-21T11:50:38.583234Z","shell.execute_reply":"2025-07-21T11:50:39.792802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, mask, label, image_id, cell_id = dataset[0]\n\nprint(f\"图像 shape: {img.shape}\")           # torch.Size([3, 256, 256])\nprint(f\"掩膜 shape: {mask.shape}\")         # torch.Size([1, 256, 256])\nprint(f\"标签向量 shape: {label.shape}\")     # torch.Size([19])\nprint(f\"image_id: {image_id} | cell_id: {cell_id}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:39.794691Z","iopub.execute_input":"2025-07-21T11:50:39.794913Z","iopub.status.idle":"2025-07-21T11:50:40.044089Z","shell.execute_reply.started":"2025-07-21T11:50:39.794892Z","shell.execute_reply":"2025-07-21T11:50:40.043341Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 定义损失函数","metadata":{}},{"cell_type":"code","source":"class SelectiveBCEWithLogitsLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.loss_fn = nn.BCEWithLogitsLoss()\n\n    def forward(self, logits, mask_onehot, label_vector):\n        \"\"\"\n        logits: (B, C, H, W)\n        mask_onehot: (B, C, H, W)\n        label_vector: (B, C)\n        \"\"\"\n        label_vector = label_vector.unsqueeze(-1).unsqueeze(-1)  # (B, C, 1, 1)\n        return self.loss_fn(logits * label_vector, mask_onehot * label_vector)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:40.045473Z","iopub.execute_input":"2025-07-21T11:50:40.045719Z","iopub.status.idle":"2025-07-21T11:50:40.050808Z","shell.execute_reply.started":"2025-07-21T11:50:40.045692Z","shell.execute_reply":"2025-07-21T11:50:40.050053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class YourModel(nn.Module):\n    def __init__(self, backbone, num_classes=19):\n        super().__init__()\n        self.backbone = backbone\n        self.num_classes = num_classes\n\n    def forward(self, x, mask=None, label_vector=None):\n        logits = self.backbone(x)  # (B, C, H, W)\n\n        if mask is not None and label_vector is not None:\n            mask_onehot = F.one_hot(mask.squeeze(1).long(), num_classes=self.num_classes)  # [B, H, W, C]\n            mask_onehot = mask_onehot.permute(0, 3, 1, 2).float()  # [B, C, H, W]\n            return logits, mask_onehot, label_vector\n        else:\n            return logits\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:40.05198Z","iopub.execute_input":"2025-07-21T11:50:40.052221Z","iopub.status.idle":"2025-07-21T11:50:40.066179Z","shell.execute_reply.started":"2025-07-21T11:50:40.052198Z","shell.execute_reply":"2025-07-21T11:50:40.065446Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 轻量级UNet","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass UNet(nn.Module):\n    def __init__(self, in_channels=3, num_classes=19, base_c=64):\n        super(UNet, self).__init__()\n\n        def conv_block(in_c, out_c):\n            return nn.Sequential(\n                nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_c),\n                nn.ReLU(inplace=True),\n                nn.Conv2d(out_c, out_c, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_c),\n                nn.ReLU(inplace=True)\n            )\n\n        self.enc1 = conv_block(in_channels, base_c)\n        self.enc2 = conv_block(base_c, base_c*2)\n        self.enc3 = conv_block(base_c*2, base_c*4)\n        self.enc4 = conv_block(base_c*4, base_c*8)\n\n        self.pool = nn.MaxPool2d(2)\n\n        self.bottleneck = conv_block(base_c*8, base_c*16)\n\n        self.up4 = nn.ConvTranspose2d(base_c*16, base_c*8, kernel_size=2, stride=2)\n        self.dec4 = conv_block(base_c*16, base_c*8)\n        self.up3 = nn.ConvTranspose2d(base_c*8, base_c*4, kernel_size=2, stride=2)\n        self.dec3 = conv_block(base_c*8, base_c*4)\n        self.up2 = nn.ConvTranspose2d(base_c*4, base_c*2, kernel_size=2, stride=2)\n        self.dec2 = conv_block(base_c*4, base_c*2)\n        self.up1 = nn.ConvTranspose2d(base_c*2, base_c, kernel_size=2, stride=2)\n        self.dec1 = conv_block(base_c*2, base_c)\n\n        self.final = nn.Conv2d(base_c, num_classes, kernel_size=1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n\n        b = self.bottleneck(self.pool(e4))\n\n        d4 = self.up4(b)\n        d4 = self.dec4(torch.cat([d4, e4], dim=1))\n        d3 = self.up3(d4)\n        d3 = self.dec3(torch.cat([d3, e3], dim=1))\n        d2 = self.up2(d3)\n        d2 = self.dec2(torch.cat([d2, e2], dim=1))\n        d1 = self.up1(d2)\n        d1 = self.dec1(torch.cat([d1, e1], dim=1))\n\n        return self.final(d1)\n# 创建模型实例用于确认构建无误\nunet = UNet(in_channels=3, num_classes=19)\nmodel = YourModel(backbone=unet, num_classes=19)  # ✅ 正确\nmodel(torch.randn(2, 3, 256, 256)).shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:40.067306Z","iopub.execute_input":"2025-07-21T11:50:40.067604Z","iopub.status.idle":"2025-07-21T11:50:41.999311Z","shell.execute_reply.started":"2025-07-21T11:50:40.067581Z","shell.execute_reply":"2025-07-21T11:50:41.998414Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 定义训练函数 测试","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, optimizer, loss_fn, device, max_steps=None):\n    model.train()\n    running_loss = 0.0\n\n    loop = tqdm(dataloader, desc=\"Training\", leave=True)\n\n    for step, (images, masks, labels) in enumerate(loop):\n        if max_steps is not None and step >= max_steps:\n            break\n\n        images, masks, labels = images.to(device), masks.to(device), labels.to(device)\n        optimizer.zero_grad()\n\n        logits, mask_onehot, label_vector = model(images, masks, labels)\n        loss = loss_fn(logits, mask_onehot, label_vector)\n\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n\n        loop.set_postfix(loss=loss.item())\n\n    return running_loss / (step + 1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:42.000605Z","iopub.execute_input":"2025-07-21T11:50:42.000944Z","iopub.status.idle":"2025-07-21T11:50:42.007468Z","shell.execute_reply.started":"2025-07-21T11:50:42.000908Z","shell.execute_reply":"2025-07-21T11:50:42.006683Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 主训练循环","metadata":{}},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = model.to(device)\nloss_fn = SelectiveBCEWithLogitsLoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\n\n\ntrain_loader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True)\n\nnum_epochs = 5\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n\n    loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n\n    for step, (images, masks, labels, _, _) in enumerate(loop):\n        images, masks, labels = images.to(device), masks.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        logits, mask_onehot, label_vector = model(images, masks, labels)\n        loss = loss_fn(logits, mask_onehot, label_vector)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        loop.set_postfix(loss=loss.item())\n\n    epoch_loss = running_loss / len(train_loader)\n    print(f\"✅ Epoch {epoch+1} completed. Average Loss: {epoch_loss:.4f}\")\n\n    # ✅ 保存模型权重\n    model_path = f\"model_epoch{epoch+1}.pth\"\n    torch.save(model.state_dict(), model_path)\n    print(f\"✅ 模型已保存: {model_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-21T11:50:42.008672Z","iopub.execute_input":"2025-07-21T11:50:42.009061Z","iopub.status.idle":"2025-07-21T11:51:08.448736Z","shell.execute_reply.started":"2025-07-21T11:50:42.009021Z","shell.execute_reply":"2025-07-21T11:51:08.447274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, mask, label, image_id, cell_id = dataset[0]\n\nimg = img.unsqueeze(0).to(device)\nmask = mask.unsqueeze(0).to(device)\n\n# 修复关键点：先转 tensor，再 unsqueeze\nif isinstance(label_vector, np.ndarray):\n    label_vector = torch.tensor(label_vector, dtype=torch.float32)\n\nlabel_vector = label_vector.unsqueeze(0).to(device)\n\nlogits, mask_onehot, label_vector = model(img, mask, label_vector)\n\nprint(\"mask shape:\", mask.shape)\nprint(\"mask_onehot shape:\", mask_onehot.shape)\nprint(\"label_vector shape:\", label_vector.shape)\nprint(\"logits shape:\", logits.shape)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# inference","metadata":{}},{"cell_type":"code","source":"# UNet 和 YourModel 定义保持不变（你已经写好了）\n# 此处略去 UNet 定义（如前所示）\n\nclass YourModel(nn.Module):\n    def __init__(self, backbone, num_classes=19):\n        super().__init__()\n        self.backbone = backbone\n        self.num_classes = num_classes\n\n    def forward(self, x, mask=None, label_vector=None):\n        logits = self.backbone(x)  # (B, C, H, W)\n\n        if mask is not None and label_vector is not None:\n            mask_onehot = F.one_hot(mask.squeeze(1).long(), num_classes=self.num_classes)  # [B, H, W, C]\n            mask_onehot = mask_onehot.permute(0, 3, 1, 2).float()  # [B, C, H, W]\n            return logits, mask_onehot, label_vector\n        else:\n            return logits\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 重建模型结构（必须跟训练时一致）\nunet = UNet(in_channels=3, num_classes=19)\nmodel = YourModel(backbone=unet, num_classes=19)\n\n# 加载模型参数\nmodel = YourModel(backbone=UNet(in_channels=3, num_classes=19))\nmodel.load_state_dict(torch.load(\"/kaggle/working/model_epoch5.pth\"))\nmodel.to(device)\nmodel.eval()  # 切换为推理模式\nprint(\"模型准备就绪\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 假设你已经定义了 YourModel 和 UNet\nbackbone = UNet(in_channels=3, num_classes=19)\nmodel = YourModel(backbone=backbone, num_classes=19)\n\n# 加载保存的权重\nmodel.load_state_dict(torch.load(\"/kaggle/working/model_epoch5.pth\", map_location=\"cpu\"))  # 或 cuda\nmodel.eval().to(device)\nprint(\"模型准备就绪\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport torch.nn.functional as F\nimport pandas as pd\nimport numpy as np\n\n# ========== Step 1: 从 dataset 随机选择一张 ==========\nidx = random.randint(0, len(dataset) - 1)\nimg, mask, label_vector, image_id, cell_id = dataset[idx]\n\n# 推理\nimg_tensor = img.unsqueeze(0).to(device)\nwith torch.no_grad():\n    pred_logits = model(img_tensor)  # [1, 19, H, W]\n    pred_probs = torch.sigmoid(pred_logits[0])  # [19, H, W]\n    image_level_probs = pred_probs.mean(dim=(1, 2))  # [19]\n\n# ========== Step 2: 获取预测向量 ==========\nthreshold = 0.5\npred_vector = (image_level_probs > threshold).int().cpu().numpy()  # [19], 0/1 向量\nselected_classes = np.where(pred_vector == 1)[0].tolist()  # 只保留置信度 > 0.5 的类别\n\n# 显示元信息\nprint(f\"\\n🖼 Image ID: {image_id}\")\nprint(f\"🔬 Cell ID: {cell_id}\")\nprint(f\"\\n📊 图像级预测 vector（置信度 > 0.5）:\")\nfor i in selected_classes:\n    print(f\"✅ Class {i:2d}: {image_level_probs[i].item():.4f}\")\n\n# ========== Step 3: 从 train.csv 读取真实标签 ==========\ntrain_df = pd.read_csv(\"/kaggle/input/hpa-single-cell-image-classification/train.csv\")\nrow = train_df[train_df['ID'] == image_id]\n\nif row.empty:\n    print(f\"\\n❌ 找不到 image_id {image_id} 在 train.csv 中的记录\")\nelse:\n    label_str = row.iloc[0][\"Label\"]\n    true_labels = list(map(int, label_str.split()))\n    true_vector = np.zeros(19, dtype=int)\n    true_vector[true_labels] = 1\n\n    # ========== Step 4: 根据预测的 selected_classes 比较 ==========\n    pred_subvector = np.zeros(19, dtype=int)\n    pred_subvector[selected_classes] = 1\n\n    true_subvector = true_vector[selected_classes]\n\n    # IoU（只对预测出的类）\n    intersection = np.logical_and(pred_subvector, true_vector).sum()\n    union = np.logical_or(pred_subvector, true_vector).sum()\n    jaccard = intersection / union if union > 0 else 0.0\n\n    # cosine 相似度\n    cosine = np.dot(pred_subvector, true_vector) / (\n        np.linalg.norm(pred_subvector) * np.linalg.norm(true_vector) + 1e-6\n    )\n\n    # 准确率（只在 selected_classes 中对比）\n    match_ratio = (pred_subvector == true_vector).sum() / 19  # 或只在 selected_classes 比例上算\n\n    # ========== Step 5: 显示对比结果 ==========\n    print(f\"\\n🎯 True Labels: {true_labels}\")\n    print(f\"✅ Jaccard (IoU): {jaccard:.4f}\")\n    print(f\"✅ Cosine similarity: {cosine:.4f}\")\n    print(f\"✅ Overlap accuracy (full vector): {match_ratio:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 可视化这些通过阈值的结构掩膜\nimport matplotlib.pyplot as plt\n\nvalid_classes = np.where(pred_vector == 1)[0]  # 获取预测为1的类别索引\n\nif len(valid_classes) == 0:\n    print(\"❌ 没有任何结构置信度 > 0.5，不进行可视化。\")\nelse:\n    num_classes = len(valid_classes)\n    plt.figure(figsize=(4 * (num_classes + 1), 4))\n\n    # 显示原图\n    img_np = img[:3].permute(1, 2, 0).cpu().numpy()\n    img_np = img_np / img_np.max()\n    plt.subplot(1, num_classes + 1, 1)\n    plt.imshow(img_np)\n    plt.title(\"Input Image\")\n\n    # 显示每个预测结构\n    for idx, class_id in enumerate(valid_classes):\n        pred_mask_cls = pred_probs[class_id].cpu().numpy()\n        plt.subplot(1, num_classes + 1, idx + 2)\n        plt.imshow(pred_mask_cls, cmap=\"viridis\")\n        plt.title(f\"Pred Mask: class {class_id}\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"version 3是在进行全新的修改之前,全部泡桐冰洁结果相对准确的version.可以随时回去调用.","metadata":{}}]}