{"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":"none","dataSources":[{"sourceId":23823,"databundleVersionId":1920183,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12124945,"sourceType":"datasetVersion","datasetId":7634750}],"dockerImageVersionId":30056,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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","execution":{"iopub.status.busy":"2025-06-11T03:51:50.420771Z","iopub.execute_input":"2025-06-11T03:51:50.421025Z","iopub.status.idle":"2025-06-11T03:51:52.114849Z","shell.execute_reply.started":"2025-06-11T03:51:50.420965Z","shell.execute_reply":"2025-06-11T03:51:52.114222Z"},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#  设定参数","metadata":{}},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-11T03:51:52.116754Z","iopub.execute_input":"2025-06-11T03:51:52.116987Z","iopub.status.idle":"2025-06-11T03:51:52.184677Z","shell.execute_reply.started":"2025-06-11T03:51:52.116965Z","shell.execute_reply":"2025-06-11T03:51:52.183797Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# check what is in weak_label","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# 路径设定\nweak_label_path = '/kaggle/input/weak-supervision-prepare-0-20-83/cell_level_weak_labels_0.83.csv'\n\n# 读取 CSV 文件\ndf = pd.read_csv(weak_label_path)\n\n# 显示前几行内容\nprint(\"文件前几行内容：\")\nprint(df.head())\n\n# 数据基本信息\nprint(\"\\n文件基本信息：\")\nprint(f\"总记录数: {len(df)}\")\nprint(f\"唯一 image_id 数: {df['image_id'].nunique()}\")\nprint(f\"唯一 cell_id 数: {df['cell_id'].nunique()} (不一定完全准确，因为 cell_id 可重复使用)\")\nprint(f\"唯一 class_id 数: {df['class_id'].nunique()}\")\n\n# 检查是否存在重复记录（即相同 image_id + cell_id + class_id 出现多次）\nduplicates = df.duplicated(subset=['image_id', 'cell_id', 'class_id'])\nprint(f\"\\n重复记录数: {duplicates.sum()}\")\n\n# 每个 cell 拥有的 class_id 数量分布\nprint(\"\\n每个细胞拥有的类别数量（前几项）：\")\ncell_class_counts = df.groupby(['image_id', 'cell_id'])['class_id'].nunique().value_counts().sort_index()\nprint(cell_class_counts.head())\n\n# 类别频率分布\nprint(\"\\n每个 class_id 出现次数（前几项）：\")\nprint(df['class_id'].value_counts().head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:05:24.679027Z","iopub.execute_input":"2025-06-10T08:05:24.679388Z","iopub.status.idle":"2025-06-10T08:05:24.732339Z","shell.execute_reply.started":"2025-06-10T08:05:24.679352Z","shell.execute_reply":"2025-06-10T08:05:24.73169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# delete missing image (cell)","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom PIL import Image\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset\n\n# 设置路径\ncell_crop_dir = '/kaggle/input/weak-supervision-prepare-0-20-83/cell_crops_0.83'\nlabel_csv_path = '/kaggle/input/weak-supervision-prepare-0-20-83/cell_level_weak_labels_0.83.csv'\n\n# 加载标签\ndf = pd.read_csv(label_csv_path)\n\n# 构建 filename 字段（与图像 crop 文件名一致）\ndf['filename'] = df.apply(\n    lambda row: f\"{row['image_id']}_class{row['class_id']}_cell{row['cell_id']}.png\", axis=1\n)\n\n# 获取实际存在的图像文件名\nall_images = set(os.listdir(cell_crop_dir))\n\n# 保留存在对应图像的标签记录\ndf_valid = df[df['filename'].isin(all_images)].reset_index(drop=True)\n\n# 识别缺失图像对应的标签记录\nmissing_files = df[~df['filename'].isin(all_images)]\n\n# 打印保留和删除的数量\ntotal = len(df)\nkept = len(df_valid)\ndropped = len(missing_files)\n\nprint(\"标签数据清洗情况：\")\nprint(f\"总标签记录数: {total}\")\nprint(f\"保留的标签记录数（文件存在）: {kept}\")\nprint(f\"删除的标签记录数（文件缺失）: {dropped}\")\nprint(f\"删除比例: {dropped / total:.2%}\")\n\n# 示例：打印前几条缺失记录\nprint(\"示例缺失记录：\")\nprint(missing_files.head())\n\n# 图像预处理（你后续训练模型时会用）\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n])\n\n# 如需保存清洗后的标签文件\ndf_valid.to_csv('/kaggle/working/cleaned_weak_labels.csv', index=False)\nprint(\"清洗后的 weak label 已保存到: /kaggle/working/cleaned_weak_labels.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:28:35.322001Z","iopub.execute_input":"2025-06-10T08:28:35.322382Z","iopub.status.idle":"2025-06-10T08:28:35.410602Z","shell.execute_reply.started":"2025-06-10T08:28:35.322346Z","shell.execute_reply":"2025-06-10T08:28:35.409803Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"all correct","metadata":{}},{"cell_type":"markdown","source":"# define PyTorch Dataset","metadata":{}},{"cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, dataframe, image_dir, transform=None):\n        self.df = dataframe\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.image_dir, row['filename'])\n        image = Image.open(img_path).convert('RGB')\n        label = int(row['class_id'])\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n# 初始化 Dataset 和 DataLoader\ndataset = CellDataset(df_valid, cell_crop_dir, transform=transform)\ntrain_loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)\n# 显示前一个样本信息\nimg, label = dataset[0]\nprint(f\"图像大小: {img.shape}, 标签: {label}\")\n\nimport numpy as np\nclass_counts = np.bincount(df_valid['class_id'])\nprint(\"每个类别的样本数:\", class_counts)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:28:43.19316Z","iopub.execute_input":"2025-06-10T08:28:43.193519Z","iopub.status.idle":"2025-06-10T08:28:43.20725Z","shell.execute_reply.started":"2025-06-10T08:28:43.193489Z","shell.execute_reply":"2025-06-10T08:28:43.206132Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"224 x 224 是使用 transforms.Resize((224, 224)) 指定的尺寸;[3] 是 RGB 通道数；表示该 cell crop 属于第 15 个类别（索引从 0 开始）","metadata":{}},{"cell_type":"markdown","source":"# define model","metadata":{}},{"cell_type":"code","source":"model = models.resnet18(pretrained=False)  # Kaggle 无法下载预训练权重\nmodel.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)\nmodel = model.to(DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:28:56.056083Z","iopub.execute_input":"2025-06-10T08:28:56.056427Z","iopub.status.idle":"2025-06-10T08:28:56.280827Z","shell.execute_reply.started":"2025-06-10T08:28:56.056396Z","shell.execute_reply":"2025-06-10T08:28:56.279681Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# loss function and optimizer","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:28:58.78983Z","iopub.execute_input":"2025-06-10T08:28:58.790148Z","iopub.status.idle":"2025-06-10T08:28:58.794932Z","shell.execute_reply.started":"2025-06-10T08:28:58.790117Z","shell.execute_reply":"2025-06-10T08:28:58.794122Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# loop train","metadata":{}},{"cell_type":"code","source":"EPOCHS = 10  # 将训练轮数设为 8\n\nmodel.train()\nfor epoch in range(EPOCHS):\n    total_loss = 0\n    correct = 0\n    total = 0\n    loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n\n    for images, labels in loop:\n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)\n\n        # forward\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        # backward\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        # metrics\n        total_loss += loss.item()\n        preds = torch.argmax(outputs, dim=1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n        loop.set_postfix(loss=loss.item(), acc=correct/total)\n\n    print(f\"Epoch {epoch+1} done. Avg loss: {total_loss/len(train_loader):.4f}, Accuracy: {correct/total:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:29:02.713857Z","iopub.execute_input":"2025-06-10T08:29:02.714182Z","iopub.status.idle":"2025-06-10T08:32:54.629054Z","shell.execute_reply.started":"2025-06-10T08:29:02.71415Z","shell.execute_reply":"2025-06-10T08:32:54.628364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" weak supervision 模型（ResNet18）已经成功收敛，并具备了较强的分类能力。\n Avg loss（平均损失）：衡量模型预测结果与真实标签之间差距的平均值。越低越好。模型从 1.57 降到了 0.19，说明模型训练非常成功。\n Accuracy（准确率）：所有样本中预测正确的比例。准确率从 49% 提高到了 93%+，说明模型基本已经学会根据 cell crop 图像判断其弱标签","metadata":{}},{"cell_type":"markdown","source":"# save model","metadata":{}},{"cell_type":"code","source":"# 保存模型权重\nmodel_path = 'resnet18_weak_supervision.pth'\ntorch.save(model.state_dict(), model_path)\nprint(f\"模型已保存为: {model_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T08:35:21.786506Z","iopub.execute_input":"2025-06-10T08:35:21.786846Z","iopub.status.idle":"2025-06-10T08:35:21.881879Z","shell.execute_reply.started":"2025-06-10T08:35:21.786816Z","shell.execute_reply":"2025-06-10T08:35:21.881093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.path.exists('/kaggle/working/resnet18_weak_supervision.pth'))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T09:17:17.840446Z","iopub.execute_input":"2025-06-10T09:17:17.840713Z","iopub.status.idle":"2025-06-10T09:17:17.844867Z","shell.execute_reply.started":"2025-06-10T09:17:17.840651Z","shell.execute_reply":"2025-06-10T09:17:17.844262Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# select confidence","metadata":{}},{"cell_type":"code","source":"import os\nimport torch.nn.functional as F\n\n# 使用与训练一致的数据集（只包含存在的图像）\nexisting_files = set(os.listdir(cell_crop_dir))\ntrain_df = df[df['filename'].isin(existing_files)].reset_index(drop=True)\n\nvalid_indices = []\nall_results = []\n\nmodel.eval()\nwith torch.no_grad():\n    for i, (images, labels) in enumerate(tqdm(train_loader, desc=\"进行预测\")):\n        batch_start = i * BATCH_SIZE\n        batch_size = images.size(0)\n\n        images = images.to(DEVICE)\n        outputs = model(images)\n        probs = F.softmax(outputs, dim=1)\n        top_probs, top_classes = torch.max(probs, dim=1)\n\n        for j in range(batch_size):\n            index_in_df = batch_start + j\n            valid_indices.append(index_in_df)\n            all_results.append({\n                'pred_class': int(top_classes[j].cpu().numpy()),\n                'confidence': float(top_probs[j].cpu().numpy())\n            })\n\n# 获取与预测一一对应的 DataFrame 行\nvalid_df = train_df.iloc[valid_indices].reset_index(drop=True)\n\n# 合并预测结果\nvalid_df['pred_class'] = [r['pred_class'] for r in all_results]\nvalid_df['confidence'] = [r['confidence'] for r in all_results]\n\n# 筛选高置信度样本\nfiltered_df = valid_df[valid_df['confidence'] >= 0.5]\n\n# 保存\nfiltered_df.to_csv('weak_supervision_labels.csv', index=False)\nprint(f\"预测完成，共保存高置信度 cell 数量: {len(filtered_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T09:35:51.26073Z","iopub.execute_input":"2025-06-05T09:35:51.261169Z","iopub.status.idle":"2025-06-05T09:38:14.180878Z","shell.execute_reply.started":"2025-06-05T09:35:51.261138Z","shell.execute_reply":"2025-06-05T09:38:14.179872Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# save the weak_supervision_labels.csv","metadata":{}},{"cell_type":"code","source":"filtered_df.to_csv('/kaggle/working/weak_supervision_labels.csv', index=False)\nprint(f\"实际保存高置信度样本数量: {len(filtered_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T10:19:21.63812Z","iopub.execute_input":"2025-06-05T10:19:21.638408Z","iopub.status.idle":"2025-06-05T10:19:21.707837Z","shell.execute_reply.started":"2025-06-05T10:19:21.638384Z","shell.execute_reply":"2025-06-05T10:19:21.707049Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"这个是我们在b5里面真正要用的，因为他们的confidence都很高，可能是因为泛化导致的，保留更有准确性","metadata":{}},{"cell_type":"markdown","source":"# what is in the csv","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# 读取文件\ndf = pd.read_csv('/kaggle/working/weak_supervision_labels.csv')\n\n# 显示前几行\nprint(\"前几行数据：\")\ndisplay(df.head())\n\n# 显示数据维度和基本统计信息\nprint(\"\\n每个预测类别的细胞数量：\")\nprint(df['pred_class'].value_counts().sort_index())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T09:38:25.778064Z","iopub.execute_input":"2025-06-05T09:38:25.778381Z","iopub.status.idle":"2025-06-05T09:38:25.829682Z","shell.execute_reply.started":"2025-06-05T09:38:25.778354Z","shell.execute_reply":"2025-06-05T09:38:25.828579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# 读取 CSV 文件\ndf = pd.read_csv('weak_supervision_labels.csv')\n\n# 计算唯一 filename 数量\nnum_unique_filenames = df['filename'].nunique()\n\nprint(f\"唯一 filename 数量为: {num_unique_filenames}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T09:44:45.796378Z","iopub.execute_input":"2025-06-05T09:44:45.796706Z","iopub.status.idle":"2025-06-05T09:44:45.823328Z","shell.execute_reply.started":"2025-06-05T09:44:45.79668Z","shell.execute_reply":"2025-06-05T09:44:45.822443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n# 读取 CSV 文件（根据你的实际路径修改）\ndf = pd.read_csv('weak_supervision_labels.csv')\n\n# 每个 filename 对应的唯一 pred_class 数量\nclass_counts_per_filename = df.groupby('filename')['pred_class'].nunique()\n\n# 统计：每种类别数（1种、2种、3种...）的 filename 有多少\ncount_distribution = class_counts_per_filename.value_counts().sort_index()\n\n# 绘图\nplt.figure(figsize=(8, 5))\ncount_distribution.plot(kind='bar')\nplt.title('Distribution of Unique Predicted Classes per Filename')\nplt.xlabel('Number of Unique Predicted Classes')\nplt.ylabel('Number of Filenames')\nplt.xticks(rotation=0)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T09:58:20.790046Z","iopub.execute_input":"2025-06-05T09:58:20.790379Z","iopub.status.idle":"2025-06-05T09:58:21.007449Z","shell.execute_reply.started":"2025-06-05T09:58:20.790345Z","shell.execute_reply":"2025-06-05T09:58:21.006692Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# if i choose top k here","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# 从已有高置信度筛选后的 CSV 中读取数据\ndf = pd.read_csv('/kaggle/working/weak_supervision_labels.csv')\n\n# 分组处理：按 filename 判断是否保留全部或 top-5 class\nfiltered_rows = []\n\nfor fname, group in df.groupby('filename'):\n    unique_classes = group['pred_class'].nunique()\n\n    if unique_classes <= 3:\n        filtered_rows.append(group)  # 保留全部\n    else:\n        # 计算每个 pred_class 的平均置信度\n        class_avg_conf = group.groupby('pred_class')['confidence'].mean()\n\n        # 获取 top-5 置信度的 pred_class\n        top5_classes = class_avg_conf.sort_values(ascending=False).head(5).index\n\n        # 只保留这 top-5 类别的 cell\n        filtered = group[group['pred_class'].isin(top5_classes)]\n        filtered_rows.append(filtered)\n\n# 合并结果 & 保存\nfiltered_df = pd.concat(filtered_rows).reset_index(drop=True)\nfiltered_df.to_csv('/kaggle/working/weak_supervision_labels_topk.csv', index=False)\nprint(f\"Top-K 过滤完成，最终保留细胞数: {len(filtered_df)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T10:15:47.741119Z","iopub.execute_input":"2025-06-05T10:15:47.741435Z","iopub.status.idle":"2025-06-05T10:18:00.062029Z","shell.execute_reply.started":"2025-06-05T10:15:47.741408Z","shell.execute_reply":"2025-06-05T10:18:00.061134Z"}},"outputs":[],"execution_count":null}]}