{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":847338,"sourceType":"datasetVersion","datasetId":255488},{"sourceId":848739,"sourceType":"datasetVersion","datasetId":251095},{"sourceId":1477289,"sourceType":"datasetVersion","datasetId":866838},{"sourceId":1983975,"sourceType":"datasetVersion","datasetId":1182793},{"sourceId":2165946,"sourceType":"datasetVersion","datasetId":1300160},{"sourceId":2182061,"sourceType":"datasetVersion","datasetId":1136931},{"sourceId":2188154,"sourceType":"datasetVersion","datasetId":1313589},{"sourceId":2194412,"sourceType":"datasetVersion","datasetId":1317588},{"sourceId":2211395,"sourceType":"datasetVersion","datasetId":1325931},{"sourceId":2211693,"sourceType":"datasetVersion","datasetId":1328208},{"sourceId":2211791,"sourceType":"datasetVersion","datasetId":1328270},{"sourceId":2211937,"sourceType":"datasetVersion","datasetId":1328360},{"sourceId":2214724,"sourceType":"datasetVersion","datasetId":1330046},{"sourceId":2388996,"sourceType":"datasetVersion","datasetId":1444318},{"sourceId":3075714,"sourceType":"datasetVersion","datasetId":849808},{"sourceId":11920265,"sourceType":"datasetVersion","datasetId":7494043},{"sourceId":12069073,"sourceType":"datasetVersion","datasetId":7596830},{"sourceId":12076621,"sourceType":"datasetVersion","datasetId":7601990},{"sourceId":12076657,"sourceType":"datasetVersion","datasetId":7602018},{"sourceId":12126734,"sourceType":"datasetVersion","datasetId":7635997},{"sourceId":12127290,"sourceType":"datasetVersion","datasetId":7636434},{"sourceId":12283545,"sourceType":"datasetVersion","datasetId":7741280},{"sourceId":3732,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":2659,"modelId":312}],"dockerImageVersionId":30096,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp -r ../input/landmark-additional-packages/rwightman_gen-efficientnet-pytorch_master/rwightman_gen-efficientnet-pytorch_master /root/.cache/torch/hub\n!cp ../input/landmark-additional-packages/tf_efficientnet_b3_aa-84b4657e.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/tf_efficientnet_b5_ra-9a3e5369.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/se_resnext50_32x4d-a260b3a4.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/resnet50d_ra2-464e36ba.pth /root/.cache/torch/hub/checkpoints/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T01:52:18.460269Z","iopub.execute_input":"2025-06-26T01:52:18.460505Z","iopub.status.idle":"2025-06-26T01:52:30.291323Z","shell.execute_reply.started":"2025-06-26T01:52:18.460449Z","shell.execute_reply":"2025-06-26T01:52:30.290281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install \"../input/landmark-additional-packages/timm-0.3.4-py3-none-any.whl\" # already in the system\n#!pip install \"../input/landmark-additional-packages/geffnet-1.0.0-py3-none-any.whl\" # already in the system","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T01:52:30.292783Z","iopub.execute_input":"2025-06-26T01:52:30.293172Z","iopub.status.idle":"2025-06-26T01:53:09.337749Z","shell.execute_reply.started":"2025-06-26T01:52:30.293126Z","shell.execute_reply":"2025-06-26T01:53:09.336961Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# new model","metadata":{}},{"cell_type":"markdown","source":"final permission: 给定一张单细胞图像，预测像素级别的结构类型（即每个像素属于哪个 class）","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nfrom PIL import Image\nimport numpy as np\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as T\n\nclass WeakSupervisionSegDataset(Dataset):\n    def __init__(self, csv_path, image_dir, mask_dir, image_size=(224, 224), num_classes=19, transforms=None):\n        self.df = pd.read_csv(csv_path)\n\n        # ✅ 确保 label_vector 字段存在\n        assert 'label_vector' in self.df.columns, \"❌ CSV 中缺少 label_vector 字段\"\n\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.image_size = image_size\n        self.num_classes = num_classes\n        self.transforms = transforms\n\n        # ✅ 分组以便按 cell_id 获取信息\n        self.cells = self.df.groupby(['image_id', 'cell_id']).apply(lambda x: x.to_dict('records')).reset_index(name='records')\n\n    def __len__(self):\n        return len(self.cells)\n\n    def find_mask_for_cell(self, image_id, cell_id):\n        \"\"\"\n        返回 mask 路径（只要找到任意 class 的一个 mask 即可作为 binary mask）\n        \"\"\"\n        folder = os.path.join(self.mask_dir, image_id)\n        if not os.path.exists(folder):\n            return None, None\n        for fname in os.listdir(folder):\n            if fname.endswith(f'cell{cell_id}.png'):\n                class_id = int(fname.split('class')[1].split('_')[0])\n                return os.path.join(folder, fname), class_id\n        return None, None\n\n    def __getitem__(self, idx):\n        records = self.cells.iloc[idx]['records']\n        image_id = records[0]['image_id']\n        cell_id = records[0]['cell_id']\n        filename = records[0]['filename']\n\n        confidences = torch.tensor([record.get('confidence', 1.0) for record in records], dtype=torch.float32)\n\n\n        # ✅ 加载图像\n        image_path = os.path.join(self.image_dir, filename)\n        image = Image.open(image_path).convert('RGB').resize(self.image_size)\n\n        # ✅ 初始化空掩膜（float 类型）\n        mask = torch.zeros((self.num_classes, *self.image_size), dtype=torch.float32)\n\n        # ✅ 读取一个 binary mask（只需要一个就能扩展多类）\n        mask_path, _ = self.find_mask_for_cell(image_id, cell_id)\n        if mask_path is not None:\n            m = Image.open(mask_path).resize(self.image_size)\n            m = (np.array(m) > 0).astype(np.float32)  # binary float mask\n\n            # ✅ 用 label_vector 填每一类的掩膜\n            label_vector = list(map(int, records[0]['label_vector'].split(',')))\n            for class_id, present in enumerate(label_vector):\n                if present == 1:\n                    mask[class_id] = torch.from_numpy(m)\n\n        # ✅ 图像转换\n        if self.transforms:\n            image = self.transforms(image)  # 如果传入的是 torchvision.transforms\n        else:\n            image = T.ToTensor()(image)  # 转成 float32 [C, H, W]\n\n        # ✅ 返回结构统一\n        return image, mask, confidences, {\n            'image_id': image_id,\n            'cell_id': cell_id,\n            'filename': filename,\n            'label_vector': torch.tensor(label_vector)  # 👈加上这个\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T01:53:09.339314Z","iopub.execute_input":"2025-06-26T01:53:09.339555Z","iopub.status.idle":"2025-06-26T01:53:11.275535Z","shell.execute_reply.started":"2025-06-26T01:53:09.339529Z","shell.execute_reply":"2025-06-26T01:53:11.274886Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"因为cell crops只是认为的复制了不同的class。所以之哟有一样的image id和cell id，就拥有一样的mask。用这个逻辑找到mask的具体文件。","metadata":{}},{"cell_type":"markdown","source":"# pos_weight","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\n\ndef compute_pos_weight_from_weak_csv(csv_path, num_classes=19, excluded_classes=[18], boost_class0=True, class0_boost_ratio=1.5, max_clip=100.0):\n    \"\"\"\n    从 weak_multilabel.csv 计算伪像素级 pos_weight。\n    每行为一个 cell，label_vector 为 '0,1,0,...' 格式。\n    \"\"\"\n    df = pd.read_csv(csv_path)\n\n    # 解析 label_vector 为 list[int]\n    df['label_list'] = df['label_vector'].apply(lambda x: list(map(int, x.split(','))))\n\n    # 累加每个类别的出现次数\n    label_counts = np.zeros(num_classes, dtype=np.int64)\n    for labels in df['label_list']:\n        for i, v in enumerate(labels):\n            if v == 1:\n                label_counts[i] += 1\n\n    print(\"📊 每个类别在 weak supervision 中出现的 cell 数量:\", label_counts)\n\n    # 频率 + pos_weight 计算\n    total = label_counts.sum() + 1e-6\n    freqs = label_counts / total\n    pos_weight = 1.0 / (freqs + 1e-6)\n\n    # 排除未出现或指定不训练的类\n    for cls in excluded_classes:\n        pos_weight[cls] = 0.0\n\n    # class_0 特别 boost\n    if boost_class0 and 0 not in excluded_classes:\n        valid_weights = pos_weight[pos_weight > 0]\n        pos_weight[0] = valid_weights.mean() * class0_boost_ratio\n\n    # 限幅\n    pos_weight = np.clip(pos_weight, 0, max_clip)\n\n    # 输出查看\n    for i, w in enumerate(pos_weight):\n        if w > 0:\n            print(f\"🧮 class_{i} → pos_weight: {w:.4f}\")\n        else:\n            print(f\"🧹 class_{i} → excluded\")\n\n    return torch.tensor(pos_weight, dtype=torch.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T01:53:11.276661Z","iopub.execute_input":"2025-06-26T01:53:11.276961Z","iopub.status.idle":"2025-06-26T01:53:11.284654Z","shell.execute_reply.started":"2025-06-26T01:53:11.276915Z","shell.execute_reply":"2025-06-26T01:53:11.283934Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 构建 Dataset + DataLoader","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader, random_split\nimport torch.nn as nn\nimport numpy as np\nimport torch\nimport torchvision.transforms as T\n\n# ✅ 数据路径\ncsv_path = '/kaggle/input/new-model-training-weak-supervision-0-003/weak_multilabel.csv'\nimage_dir = '/kaggle/input/weak-supervision-0-003/cell_crops'\nmask_dir = '/kaggle/input/mmdetec-0-003/masks_thr05/masks_thr05'\n\n# ✅ 图像预处理\ndefault_transforms = T.Compose([\n    T.ToTensor(),\n    T.Normalize(mean=[0.5]*3, std=[0.5]*3)\n])\n\n# ✅ 构建数据集\ndataset = WeakSupervisionSegDataset(csv_path, image_dir, mask_dir, transforms=default_transforms)\n\n# ✅ 划分训练 / 验证集\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\n# ✅ DataLoader\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=2, pin_memory=True)\n\n# ✅ 指定设备\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ✅ 定义 loss（不再使用 pos_weight）\nclass SelectiveBCEOnlyLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')  # 无 pos_weight\n\n    def forward(self, pred, target, label_vector):\n        B, C, H, W = pred.shape\n        loss = self.bce(pred, target)  # [B, 19, H, W]\n        label_mask = label_vector.unsqueeze(-1).unsqueeze(-1)  # [B, 19, 1, 1]\n        masked_loss = loss * label_mask\n        bce_mean = masked_loss.sum() / (label_mask.sum() * H * W + 1e-6)\n        return bce_mean\n\n# ✅ 初始化 loss 函数\nloss_fn = SelectiveBCEOnlyLoss().to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:37:16.159659Z","iopub.execute_input":"2025-06-26T03:37:16.159998Z","iopub.status.idle":"2025-06-26T03:37:16.361133Z","shell.execute_reply.started":"2025-06-26T03:37:16.159966Z","shell.execute_reply":"2025-06-26T03:37:16.360484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\n# 🔍 想检查的目标\ntarget_image_id = '5b88d5e8-bb99-11e8-b2b9-ac1f6b6435d0'\ntarget_cell_id = 0\n\n# ✅ 遍历 dataset 找匹配项\nfor i in range(len(dataset)):\n    _, _, _, meta = dataset[i]\n    if meta['image_id'] == target_image_id and meta['cell_id'] == target_cell_id:\n        print(f\"✅ 找到匹配样本！文件名：{meta['filename']}\")\n        \n        # 可视化\n        img_tensor, mask_tensor, _, _ = dataset[i]\n        img_np = img_tensor.permute(1, 2, 0).numpy()\n        img_np = (img_np * 0.5 + 0.5).clip(0, 1)  # 反归一化\n        \n        plt.figure(figsize=(4, 4))\n        plt.imshow(img_np)\n        plt.title(f\"{meta['filename']}\")\n        plt.axis('off')\n        plt.show()\n        break\nelse:\n    print(\"❌ 没找到对应 image_id + cell_id 的样本！\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:37:23.573151Z","iopub.execute_input":"2025-06-26T03:37:23.573447Z","iopub.status.idle":"2025-06-26T03:37:23.776105Z","shell.execute_reply.started":"2025-06-26T03:37:23.573421Z","shell.execute_reply":"2025-06-26T03:37:23.775236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 先从 dataset 中拿一个样本，例如 index = 0\nsample = dataset[0]\n\nimage = sample[0]  # Tensor\nmask = sample[1]   # [19, H, W]\nconf = sample[2]   # confidences（可选）\nmeta = sample[3]   # 包含 image_id, cell_id, filename, label_vector\n\n# 获取标签向量\nlabel_vector = meta['label_vector']  # Tensor [19]\nlabel_vector_np = label_vector.numpy()\n\n# 找出哪些结构类型为正例（即值为1）\npresent_classes = [i for i, v in enumerate(label_vector_np) if v == 1]\n\n# 打印信息\nprint(f\"📌 样本文件: {meta['filename']}\")\nprint(f\"📌 image_id: {meta['image_id']} | cell_id: {meta['cell_id']}\")\nprint(f\"✅ 被告知的结构类型（cell types）是: {present_classes}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:43:34.92032Z","iopub.execute_input":"2025-06-26T03:43:34.920609Z","iopub.status.idle":"2025-06-26T03:43:34.975521Z","shell.execute_reply.started":"2025-06-26T03:43:34.920586Z","shell.execute_reply":"2025-06-26T03:43:34.974937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport torch\n\n# ✅ 添加 efficientnet_pytorch 的源码路径（用于 EfficientNet encoder）\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master')\n\n# ✅ 添加 segmentation_models_pytorch 的源码路径\nsys.path.append('/kaggle/input/segmentation-models-pytorch')\n\n# ✅ 添加 pretrainedmodels 的源码路径（用于 ResNet encoder 等）\nsys.path.append('/kaggle/input/pytorch-pretrained-models/repository/pretrained-models.pytorch-master')\n\n# ✅ 开始导入模块\nfrom efficientnet_pytorch import EfficientNet\nimport segmentation_models_pytorch as smp\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:28.24085Z","iopub.execute_input":"2025-06-26T02:01:28.241218Z","iopub.status.idle":"2025-06-26T02:01:28.245712Z","shell.execute_reply.started":"2025-06-26T02:01:28.241173Z","shell.execute_reply":"2025-06-26T02:01:28.244986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nimport torch\nimport torch.nn as nn\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ✅ 1. 创建 timm 特征提取模型（features_only 返回每个 stage 的特征）\nencoder = timm.create_model(\n    \"tf_efficientnet_b3\",\n    features_only=True,\n    pretrained=False\n)\n\n# ✅ 2. 加载预训练权重\nstate_dict = torch.load(\"/kaggle/input/tf-efficientnet-b3-weights/tf_efficientnet_b3_aa-84b4657e.pth\")\nstate_dict.pop(\"_fc.weight\", None)\nstate_dict.pop(\"_fc.bias\", None)\nencoder.load_state_dict(state_dict, strict=False)\n\n# ✅ 3. 包装成 SMP 可识别的 Encoder（这里可参考 SMP 的 encoder 格式）\nclass CustomEncoder(nn.Module):\n    def __init__(self, encoder):\n        super().__init__()\n        self.encoder = encoder\n\n    def forward(self, x):\n        features = self.encoder(x)  # list of feature maps\n        return features\n\ncustom_encoder = CustomEncoder(encoder).to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:28.247428Z","iopub.execute_input":"2025-06-26T02:01:28.247746Z","iopub.status.idle":"2025-06-26T02:01:28.575365Z","shell.execute_reply.started":"2025-06-26T02:01:28.247714Z","shell.execute_reply":"2025-06-26T02:01:28.57469Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x = torch.randn(1, 3, 224, 224).to(device)\nfeatures = custom_encoder(x)\nfor i, f in enumerate(features):\n    print(f\"Stage {i+1}: {f.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:28.57663Z","iopub.execute_input":"2025-06-26T02:01:28.576849Z","iopub.status.idle":"2025-06-26T02:01:28.640146Z","shell.execute_reply.started":"2025-06-26T02:01:28.576827Z","shell.execute_reply":"2025-06-26T02:01:28.639208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass CustomUnet(nn.Module):\n    def __init__(self, encoder, num_classes=19):\n        super().__init__()\n        self.encoder = encoder\n        self.decoder = smp.unet.decoder.UnetDecoder(\n            encoder_channels=[24, 32, 48, 136, 384],  # ← 根据 EfficientNet-B3 输出通道设置\n            decoder_channels=[256, 128, 64, 32, 32],\n            n_blocks=5,\n            use_batchnorm=True,\n            center=False,\n        )\n        self.segmentation_head = nn.Conv2d(32, num_classes, kernel_size=1)\n\n    def forward(self, x):\n        features = self.encoder(x)\n\n        # ✅ 如不需要调试可注释以下 print\n        # print(\"🔍 Encoder 输出特征图：\")\n        # for i, f in enumerate(features):\n        #     print(f\"Stage {i+1}: {f.shape}\")\n\n        x = self.decoder(*features)\n        # print(f\"✅ Decoder 输出形状: {x.shape}\")\n\n        x = self.segmentation_head(x)  # [B, 19, 112, 112]\n\n        # ✅ 上采样回原始图像大小\n        x = F.interpolate(x, size=(224, 224), mode='bilinear', align_corners=False)\n\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:28.642252Z","iopub.execute_input":"2025-06-26T02:01:28.642646Z","iopub.status.idle":"2025-06-26T02:01:28.649952Z","shell.execute_reply.started":"2025-06-26T02:01:28.642604Z","shell.execute_reply":"2025-06-26T02:01:28.648985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\nclass EfficientB3Classifier(nn.Module):\n    def __init__(self, num_classes=19):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"tf_efficientnet_b3\",\n            features_only=True,\n            pretrained=True\n        )\n        # EfficientNet-B3 最后一层输出是 stage5，对应的输出通道是 384（来自 timm 文档）\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.classifier = nn.Linear(384, num_classes)\n\n    def forward(self, x):\n        feats = self.backbone(x)           # list of feature maps\n        last_feat = feats[-1]              # 取最后一层 [B, 384, H, W]\n        pooled = self.pool(last_feat).squeeze(-1).squeeze(-1)  # [B, 384]\n        out = self.classifier(pooled)      # [B, 19]\n        return out\n\n\n# ✅ 实例化模型\nmodel = EfficientB3Classifier(num_classes=19)\nx = torch.randn(8, 3, 224, 224)  # 假设输入为 cell crop\nout = model(x)\nprint(out.shape)  # ✅ 应为 [8, 19]\nprint(\"✅ 模型结构初始化完成\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:28.651227Z","iopub.execute_input":"2025-06-26T02:01:28.651481Z","iopub.status.idle":"2025-06-26T02:01:30.276017Z","shell.execute_reply.started":"2025-06-26T02:01:28.651456Z","shell.execute_reply":"2025-06-26T02:01:30.27519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass PixelWiseWeightedBCEDiceLoss(nn.Module):\n    def __init__(self, pos_weight):\n        super().__init__()\n        self.pos_weight = pos_weight.view(1, -1, 1, 1)  # [1, 19, 1, 1]\n        self.bce_loss = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, pred, target, label_vector):\n        # pred, target: [B, 19, H, W]\n        # label_vector: [B, 19]\n\n        B, C, H, W = pred.shape\n\n        # step 1: BCE loss + 权重\n        bce = self.bce_loss(pred, target)\n        weighted_bce = bce * self.pos_weight.to(pred.device)\n\n        # step 2: mask 出有效类别通道：label_vector 控制哪些结构参与训练\n        class_mask = label_vector.unsqueeze(-1).unsqueeze(-1)  # [B, 19, 1, 1]\n        masked_bce = weighted_bce * class_mask\n        bce_mean = masked_bce.sum() / (class_mask.sum() * H * W + 1e-6)\n\n        # step 3: Dice loss 只对有效通道计算\n        pred_sigmoid = torch.sigmoid(pred)\n        inter = (pred_sigmoid * target * class_mask).sum(dim=(2,3))\n        union = (pred_sigmoid * class_mask).sum(dim=(2,3)) + (target * class_mask).sum(dim=(2,3))\n        dice = 1 - ((2 * inter + 1.) / (union + 1.))\n\n        dice_mean = (dice * class_mask.squeeze(-1).squeeze(-1)).sum() / (class_mask.sum() + 1e-6)\n\n        return bce_mean + dice_mean\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:01:30.277261Z","iopub.execute_input":"2025-06-26T02:01:30.27759Z","iopub.status.idle":"2025-06-26T02:01:30.284761Z","shell.execute_reply.started":"2025-06-26T02:01:30.277554Z","shell.execute_reply":"2025-06-26T02:01:30.284016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 正确版本：从 train_loader 拿 batch，并分析 mask\nimages, masks, _, _ = next(iter(train_loader))\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = CustomUnet(custom_encoder, num_classes=19).to(device)\ncriterion = SelectiveBCEOnlyLoss().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:02:50.396745Z","iopub.execute_input":"2025-06-26T02:02:50.397117Z","iopub.status.idle":"2025-06-26T02:02:51.851468Z","shell.execute_reply.started":"2025-06-26T02:02:50.397083Z","shell.execute_reply":"2025-06-26T02:02:51.850373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 👇 测试 batch 图像\nimages, _, _, _ = next(iter(train_loader))\nimages = images.to(device)\n\n# 👇 分段运行 encoder 和 decoder\nwith torch.no_grad():\n    features = model.encoder(images)         # List of feature maps\n    for i, f in enumerate(features):\n        print(f\"Stage {i+1} output: {f.shape}\")  # 每一层 encoder 的输出 shape\n\n    decoded = model.decoder(*features)       # 运行 decoder\n    print(f\"✅ Decoder output shape: {decoded.shape}\")  # 我们最关心的这句\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:02:51.853128Z","iopub.execute_input":"2025-06-26T02:02:51.853423Z","iopub.status.idle":"2025-06-26T02:02:53.258362Z","shell.execute_reply.started":"2025-06-26T02:02:51.853392Z","shell.execute_reply":"2025-06-26T02:02:53.257417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch\n\nbest_val_loss = float('inf')\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2)\n\nfor epoch in range(10):\n    model.train()\n    running_loss = 0.0\n    train_loader_tqdm = tqdm(train_loader, desc=f\"[Train] Epoch {epoch+1}\", leave=False)\n\n    for images, masks, _, metas in train_loader_tqdm:\n        images = images.to(device)\n        masks = masks.to(device).float()\n        label_vectors = metas['label_vector'].to(device).float()\n\n        outputs = model(images)  # [B, 19, H, W]\n\n        # ✅ 只对存在的类别计算 BCE loss\n        loss = loss_fn(outputs, masks, label_vectors)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        train_loader_tqdm.set_postfix(loss=loss.item())\n\n    avg_train_loss = running_loss / len(train_loader)\n\n    # ✅ 验证阶段\n    model.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for images, masks, _, metas in val_loader:\n            images = images.to(device)\n            masks = masks.to(device).float()\n            label_vectors = metas['label_vector'].to(device).float()\n\n            outputs = model(images)\n            loss = loss_fn(outputs, masks, label_vectors)\n            val_loss += loss.item()\n\n    avg_val_loss = val_loss / len(val_loader)\n    scheduler.step(avg_val_loss)\n\n    if avg_val_loss < best_val_loss:\n        best_val_loss = avg_val_loss\n        torch.save(model.state_dict(), f\"/kaggle/working/best_model_epoch{epoch+1}.pth\")\n        print(f\"📌 Best model saved at epoch {epoch+1} with val loss {avg_val_loss:.4f}\")\n\n    print(f\"✅ Epoch {epoch+1} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:02:53.260379Z","iopub.execute_input":"2025-06-26T02:02:53.260643Z","iopub.status.idle":"2025-06-26T02:08:48.391548Z","shell.execute_reply.started":"2025-06-26T02:02:53.260614Z","shell.execute_reply":"2025-06-26T02:08:48.390525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"✅ Dataset 总样本数:\", len(dataset))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:08:48.393193Z","iopub.execute_input":"2025-06-26T02:08:48.393465Z","iopub.status.idle":"2025-06-26T02:08:48.397814Z","shell.execute_reply.started":"2025-06-26T02:08:48.393433Z","shell.execute_reply":"2025-06-26T02:08:48.397039Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# next","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\n\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\nmodel.eval()\ncount = 0\n\nwith torch.no_grad():\n    for batch_idx, (images, masks, confidences, metas) in enumerate(val_loader):\n        images = images.to(device)\n        outputs = model(images)\n        probs = torch.sigmoid(outputs)  # [B, 19, H, W]\n        \n        for i in range(len(images)):\n            image_id = metas['image_id'][i]\n            if image_id != target_image_id:\n                continue\n            \n            count += 1\n            meta = {\n                'image_id': image_id,\n                'cell_id': metas['cell_id'][i],\n                'filename': metas['filename'][i]\n            }\n\n            image_np = images[i].cpu().permute(1, 2, 0).numpy()\n            mask_np = masks[i].cpu().numpy()\n            prob_np = probs[i].cpu().numpy()\n\n            max_conf_per_class = prob_np.reshape(19, -1).max(axis=1)\n            topk = 3\n            top_classes = max_conf_per_class.argsort()[-topk:][::-1]\n            top_confs = max_conf_per_class[top_classes]\n\n            print(f\"\\n📌 image_id: {meta['image_id']}, cell_id: {meta['cell_id']}\")\n            print(f\"📊 全部19类的最大pixel-level置信度向量：\\n{max_conf_per_class}\")\n            for cls, conf in zip(top_classes, top_confs):\n                print(f\"  🔹 Class {cls}: max pixel confidence = {conf:.4f}\")\n\n            fig, axs = plt.subplots(1, topk + 1, figsize=(15, 5))\n            axs[0].imshow(image_np)\n            axs[0].set_title(\"Original Image\")\n            axs[0].axis('off')\n\n            for j, cls in enumerate(top_classes):\n                axs[j + 1].imshow(image_np)  # 底图\n                axs[j + 1].imshow(prob_np[cls], cmap='inferno', alpha=0.6)  # 叠加热力图\n                axs[j + 1].set_title(f'Class {cls} (conf: {top_confs[j]:.2f})')\n                axs[j + 1].axis('off')\n\n            plt.tight_layout()\n            plt.show()\n\nprint(f\"\\n✅ 共预测出 {count} 个 cell 属于图像 {target_image_id}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:08:48.398953Z","iopub.execute_input":"2025-06-26T02:08:48.399266Z","iopub.status.idle":"2025-06-26T02:08:56.909213Z","shell.execute_reply.started":"2025-06-26T02:08:48.399241Z","shell.execute_reply":"2025-06-26T02:08:56.908387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取原始 train.csv\ndf_train = pd.read_csv('/kaggle/input/hpa-single-cell-image-classification/train.csv')\n\n# 查找对应 image_id 的标签\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\nlabels = df_train.loc[df_train['ID'] == target_image_id, 'Label'].values\n\nif len(labels) > 0:\n    print(f\"🎯 image_id: {target_image_id} 原始标签字段为: {labels[0]}\")\n    \n    # 自动判断是 ' ' 还是 '|' 分隔符\n    if '|' in labels[0]:\n        class_ids = list(map(int, labels[0].split('|')))\n    else:\n        class_ids = list(map(int, labels[0].split()))\n\n    print(f\"✅ 对应的细胞结构类型为: {class_ids}\")\nelse:\n    print(f\"❌ image_id: {target_image_id} 未在 train.csv 中找到\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:56:26.361745Z","iopub.execute_input":"2025-06-26T02:56:26.362224Z","iopub.status.idle":"2025-06-26T02:56:26.391598Z","shell.execute_reply.started":"2025-06-26T02:56:26.362182Z","shell.execute_reply":"2025-06-26T02:56:26.390931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom collections import Counter\n\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\nmodel.eval()\npredicted_class_counter = Counter()\ncount = 0\n\nwith torch.no_grad():\n    for batch_idx, (images, masks, confidences, metas) in enumerate(val_loader):\n        images = images.to(device)\n        outputs = model(images)\n        probs = torch.sigmoid(outputs)\n\n        for i in range(len(images)):\n            image_id = metas['image_id'][i]\n            if image_id != target_image_id:\n                continue\n\n            count += 1\n            prob_np = probs[i].cpu().numpy()\n            max_conf_per_class = prob_np.reshape(19, -1).max(axis=1)\n            topk = 3\n            top_classes = max_conf_per_class.argsort()[-topk:][::-1]\n\n            for cls in top_classes:\n                predicted_class_counter[cls] += 1\n\nprint(f\"\\n✅ image_id = {target_image_id} 中共有 {count} 个 cell 被预测\")\nprint(f\"📊 Top3 类别预测频次统计如下：\")\nfor cls, freq in predicted_class_counter.most_common():\n    print(f\"🔹 class {cls}: 出现次数 = {freq}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:08:56.953975Z","iopub.execute_input":"2025-06-26T02:08:56.95431Z","iopub.status.idle":"2025-06-26T02:09:03.689323Z","shell.execute_reply.started":"2025-06-26T02:08:56.954274Z","shell.execute_reply":"2025-06-26T02:09:03.688361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 文件路径\ncsv_path = '/kaggle/input/new-model-training-weak-supervision-0-003/weak_multilabel.csv'\n\n# 加载 CSV\ndf = pd.read_csv(csv_path)\n\n# 指定要查询的 image_id\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\n\n# 过滤出该 image_id 的所有 cell\nsub_df = df[df['image_id'] == target_image_id].copy()\n\n# 解析 label_vector 为 list[int]\nsub_df['label_list'] = sub_df['label_vector'].apply(lambda x: list(map(int, x.split(','))))\n\n# 累加所有 cell 的标签，判断哪些 class 至少为 1\nimport numpy as np\nlabel_matrix = np.array(sub_df['label_list'].tolist())\nclass_presence = (label_matrix.sum(axis=0) > 0).astype(int)\n\n# 输出结果\nprint(f\"📌 image_id = {target_image_id} 被预测为正例的结构类型有:\")\nfor i, present in enumerate(class_presence):\n    if present:\n        print(f\"✅ class {i}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:09:03.691552Z","iopub.execute_input":"2025-06-26T02:09:03.691808Z","iopub.status.idle":"2025-06-26T02:09:03.709221Z","shell.execute_reply.started":"2025-06-26T02:09:03.691778Z","shell.execute_reply":"2025-06-26T02:09:03.708454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 加载 CSV\ncsv_path = '/kaggle/input/new-model-training-weak-supervision-0-003/weak_multilabel.csv'\ndf = pd.read_csv(csv_path)\n\n# 筛选特定 image_id 的行\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\nsubset = df[df['image_id'] == target_image_id]\n\n# 获取 unique cell_id 数量\nunique_cell_ids = subset['cell_id'].nunique()\nprint(f\"🧪 图像 {target_image_id} 中共有 {unique_cell_ids} 个唯一 cell_id\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:09:03.710564Z","iopub.execute_input":"2025-06-26T02:09:03.710908Z","iopub.status.idle":"2025-06-26T02:09:03.721168Z","shell.execute_reply.started":"2025-06-26T02:09:03.710855Z","shell.execute_reply":"2025-06-26T02:09:03.720336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 假设你已经有 val_dataset 和 image_id\ntarget_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\nval_count = 0\nfor i in range(len(val_dataset)):\n    item = val_dataset[i]\n    if item[3]['image_id'] == target_id:\n        val_count += 1\n\nprint(f\"📊 在验证集中，图像 {target_id} 有 {val_count} 个 cell 被用于验证\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:09:03.722348Z","iopub.execute_input":"2025-06-26T02:09:03.722668Z","iopub.status.idle":"2025-06-26T02:09:13.993824Z","shell.execute_reply.started":"2025-06-26T02:09:03.722635Z","shell.execute_reply":"2025-06-26T02:09:13.992966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\ntrain_count = 0\nfor i in range(len(train_dataset)):\n    item = train_dataset[i]\n    if item[3]['image_id'] == target_id:\n        train_count += 1\n\nprint(f\"📊 在训练集中，图像 {target_id} 有 {train_count} 个 cell 被用于训练\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:09:13.995177Z","iopub.execute_input":"2025-06-26T02:09:13.995529Z","iopub.status.idle":"2025-06-26T02:09:55.039497Z","shell.execute_reply.started":"2025-06-26T02:09:13.995492Z","shell.execute_reply":"2025-06-26T02:09:55.038568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 载入 CSV 文件\ncsv_path = '/kaggle/input/new-model-training-weak-supervision-0-003/weak_multilabel.csv'\ndf = pd.read_csv(csv_path)\n\n# 设定目标 image_id\ntarget_image_id = '5b931256-bb99-11e8-b2b9-ac1f6b6435d0'\n\n# 筛选出该图像的所有记录\ndf_target = df[df['image_id'] == target_image_id]\n\n# 唯一的 filename 数量\nunique_filenames = df_target['filename'].nunique()\n\n# 唯一的 (image_id, cell_id) 数量\nunique_cells = df_target[['image_id', 'cell_id']].drop_duplicates().shape[0]\n\n# 输出结果\nprint(f\"📁 图像 {target_image_id} 中有 {unique_filenames} 个唯一 filename\")\nprint(f\"🧫 图像 {target_image_id} 中有 {unique_cells} 个唯一 (image_id, cell_id) 组合（即不同细胞）\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:09:55.040876Z","iopub.execute_input":"2025-06-26T02:09:55.041275Z","iopub.status.idle":"2025-06-26T02:09:55.057424Z","shell.execute_reply.started":"2025-06-26T02:09:55.041234Z","shell.execute_reply":"2025-06-26T02:09:55.056661Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 修改记录","metadata":{}},{"cell_type":"markdown","source":"第二次修改：引入 label_vector 控制监督通道：\n每个细胞只计算其真实存在结构的类别损失，避免误导模型学习不存在的结构。\n\n自定义 Loss 函数（BCE + Dice）加强有效通道训练：\n同时考虑类别不平衡（pos_weight）与像素重叠程度（Dice loss）。\n\n训练 loop 中增加了 label_vector 的传入与使用：\n在训练和验证阶段，loss 仅在有效结构上进行计算，更精确更高效。","metadata":{}},{"cell_type":"markdown","source":"第三次修改打算：第三次修改打算：取消 class weight，仅用 selective supervision（即 label_vector 屏蔽非目标类别）+ Dice loss 来训练。","metadata":{}},{"cell_type":"markdown","source":"结果到这一步发现，数据的选择还是dataset的搭建都是正确的。只能怀疑是模型的问题第四次修改：检查一下我现在mask和lable是不是对齐的.","metadata":{}},{"cell_type":"markdown","source":"# debug part","metadata":{}},{"cell_type":"code","source":"for i in range(10):\n    img, mask, _, meta = dataset[i]\n    label_vec = meta['label_vector']\n    for c in range(19):\n        if label_vec[c] == 1:\n            print(f\"cell {i} 结构 {c} | mask sum: {mask[c].sum()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T02:15:49.767865Z","iopub.execute_input":"2025-06-26T02:15:49.768201Z","iopub.status.idle":"2025-06-26T02:15:50.234033Z","shell.execute_reply.started":"2025-06-26T02:15:49.768174Z","shell.execute_reply":"2025-06-26T02:15:50.233214Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 配对cell crops 还有正确的mask（crops clean版）","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# 路径设定\nmask_root = '/kaggle/input/mmdetec-0-003/masks_thr05/masks_thr05'\ncell_crops_root = '/kaggle/input/weak-supervision-0-003/cell_crops'\n\nmatched_files = []\n\n# 遍历每个 image_id 文件夹\nfor image_id in os.listdir(mask_root):\n    mask_dir = os.path.join(mask_root, image_id)\n    if not os.path.isdir(mask_dir):\n        continue\n    \n    for fname in os.listdir(mask_dir):\n        # 提取结构类别和 cell_id\n        if fname.startswith(\"class\") and \"_cell\" in fname:\n            parts = fname.replace(\".png\", \"\").split(\"_\")\n            class_id = parts[0].replace(\"class\", \"\")\n            cell_id = parts[1].replace(\"cell\", \"\")\n            \n            # 生成对应的 crop 文件名\n            crop_fname = f\"{image_id}_class{class_id}_cell{cell_id}.png\"\n            crop_path = os.path.join(cell_crops_root, crop_fname)\n            \n            if os.path.exists(crop_path):\n                matched_files.append({\n                    'image_id': image_id,\n                    'class_id': int(class_id),\n                    'cell_id': int(cell_id),\n                    'crop_filename': crop_fname\n                })\n\n# 保存为 CSV\ndf = pd.DataFrame(matched_files)\ndf.to_csv('/kaggle/working/matched_cell_crops_from_masks.csv', index=False)\nprint(f\"✅ 共匹配成功 {len(df)} 个 cell crops 文件，并已保存为 CSV。\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:06:37.077633Z","iopub.execute_input":"2025-06-26T03:06:37.077932Z","iopub.status.idle":"2025-06-26T03:06:38.525799Z","shell.execute_reply.started":"2025-06-26T03:06:37.077905Z","shell.execute_reply":"2025-06-26T03:06:38.524966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 文件路径\ncsv_path = '/kaggle/input/new-model-training-weak-supervision-0-003/weak_multilabel.csv'\n\n# 加载 CSV\ndf = pd.read_csv(csv_path)\n\n# 指定要查询的 image_id\ntarget_image_id = '5b88d5e8-bb99-11e8-b2b9-ac1f6b6435d0'\n\n# 过滤出该 image_id 的所有 cell\nsub_df = df[df['image_id'] == target_image_id].copy()\n\n# 解析 label_vector 为 list[int]\nsub_df['label_list'] = sub_df['label_vector'].apply(lambda x: list(map(int, x.split(','))))\n\n# 累加所有 cell 的标签，判断哪些 class 至少为 1\nimport numpy as np\nlabel_matrix = np.array(sub_df['label_list'].tolist())\nclass_presence = (label_matrix.sum(axis=0) > 0).astype(int)\n\n# 输出结果\nprint(f\"📌 image_id = {target_image_id} 被预测为正例的结构类型有:\")\nfor i, present in enumerate(class_presence):\n    if present:\n        print(f\"✅ class {i}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:39:41.436452Z","iopub.execute_input":"2025-06-26T03:39:41.436846Z","iopub.status.idle":"2025-06-26T03:39:41.45432Z","shell.execute_reply.started":"2025-06-26T03:39:41.436812Z","shell.execute_reply":"2025-06-26T03:39:41.45323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取原始 train.csv\ndf_train = pd.read_csv('/kaggle/input/hpa-single-cell-image-classification/train.csv')\n\n# 查找对应 image_id 的标签\ntarget_image_id = '5b88d5e8-bb99-11e8-b2b9-ac1f6b6435d0'\nlabels = df_train.loc[df_train['ID'] == target_image_id, 'Label'].values\n\nif len(labels) > 0:\n    print(f\"🎯 image_id: {target_image_id} 原始标签字段为: {labels[0]}\")\n    \n    # 自动判断是 ' ' 还是 '|' 分隔符\n    if '|' in labels[0]:\n        class_ids = list(map(int, labels[0].split('|')))\n    else:\n        class_ids = list(map(int, labels[0].split()))\n\n    print(f\"✅ 对应的细胞结构类型为: {class_ids}\")\nelse:\n    print(f\"❌ image_id: {target_image_id} 未在 train.csv 中找到\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T03:39:29.618413Z","iopub.execute_input":"2025-06-26T03:39:29.618695Z","iopub.status.idle":"2025-06-26T03:39:29.648014Z","shell.execute_reply.started":"2025-06-26T03:39:29.618672Z","shell.execute_reply":"2025-06-26T03:39:29.647269Z"}},"outputs":[],"execution_count":null}]}