{"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":"none","dataSources":[{"sourceId":23823,"databundleVersionId":1920183,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":1983975,"sourceType":"datasetVersion","datasetId":1182793},{"sourceId":2153167,"sourceType":"datasetVersion","datasetId":1229677},{"sourceId":2165946,"sourceType":"datasetVersion","datasetId":1300160},{"sourceId":2167436,"sourceType":"datasetVersion","datasetId":1149113},{"sourceId":2182061,"sourceType":"datasetVersion","datasetId":1136931},{"sourceId":2185494,"sourceType":"datasetVersion","datasetId":1136885},{"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":2215761,"sourceType":"datasetVersion","datasetId":1330639},{"sourceId":3075714,"sourceType":"datasetVersion","datasetId":849808},{"sourceId":11920265,"sourceType":"datasetVersion","datasetId":7494043},{"sourceId":54737821,"sourceType":"kernelVersion"},{"sourceId":62498339,"sourceType":"kernelVersion"},{"sourceId":241381972,"sourceType":"kernelVersion"}],"dockerImageVersionId":30097,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-05-26T09:41:44.638177Z","iopub.execute_input":"2025-05-26T09:41:44.638704Z","iopub.status.idle":"2025-05-26T09:41:53.847182Z","shell.execute_reply.started":"2025-05-26T09:41:44.638667Z","shell.execute_reply":"2025-05-26T09:41:53.845464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 安装关键包，从你拥有的 input 路径中提取\n\n!pip install \"../input/landmark-additional-packages/EfficientNet-PyTorch/EfficientNet-PyTorch-master\" # sucess\n!pip install \"../input/landmark-additional-packages/pycocotools-2.0.2/dist/pycocotools-2.0.2.tar\"  #success","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:41:53.849404Z","iopub.execute_input":"2025-05-26T09:41:53.849727Z","iopub.status.idle":"2025-05-26T09:43:18.914028Z","shell.execute_reply.started":"2025-05-26T09:41:53.849688Z","shell.execute_reply":"2025-05-26T09:43:18.912961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install \"../input/landmark-additional-packages/faiss_gpu-1.7.0-cp37-cp37m-manylinux2014_x86_64.whl\" #may fail\n!pip install \"../input/landmark-additional-packages/pytorch_zoo-master\" # mayfail","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:43:18.916811Z","iopub.execute_input":"2025-05-26T09:43:18.917249Z","iopub.status.idle":"2025-05-26T09:44:27.519102Z","shell.execute_reply.started":"2025-05-26T09:43:18.917192Z","shell.execute_reply":"2025-05-26T09:44:27.517739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls ../input/landmark-additional-packages/\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:44:27.521532Z","iopub.execute_input":"2025-05-26T09:44:27.521958Z","iopub.status.idle":"2025-05-26T09:44:28.762173Z","shell.execute_reply.started":"2025-05-26T09:44:27.521905Z","shell.execute_reply":"2025-05-26T09:44:28.761091Z"}},"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-05-26T09:44:28.763673Z","iopub.execute_input":"2025-05-26T09:44:28.763961Z","iopub.status.idle":"2025-05-26T09:45:45.843914Z","shell.execute_reply.started":"2025-05-26T09:44:28.763927Z","shell.execute_reply":"2025-05-26T09:45:45.842687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install \"../input/landmark-additional-packages/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\" # success","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:45:45.845586Z","iopub.execute_input":"2025-05-26T09:45:45.845925Z","iopub.status.idle":"2025-05-26T09:46:25.848407Z","shell.execute_reply.started":"2025-05-26T09:45:45.845889Z","shell.execute_reply":"2025-05-26T09:46:25.847282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# the raw data has\n#!pip install -q ../input/landmark-additional-packages/timm-0.3.4-py3-none-any.whl\n#!pip install -q ../input/landmark-additional-packages/geffnet-1.0.0-py3-none-any.whl\n#!pip install -q ../input/landmark-additional-packages/EfficientNet-PyTorch/EfficientNet-PyTorch-master\n#!pip install -q ../input/landmark-additional-packages/pycocotools-2.0.2/dist/pycocotools-2.0.2.tar\n#!pip install -q ../input/landmark-additional-packages/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:46:25.849935Z","iopub.execute_input":"2025-05-26T09:46:25.850285Z","iopub.status.idle":"2025-05-26T09:46:25.854214Z","shell.execute_reply.started":"2025-05-26T09:46:25.85025Z","shell.execute_reply":"2025-05-26T09:46:25.853181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# the raw data has\n#!pip install \"/kaggle/input/hpamisc/pytorch_zoo-master\"\n#!pip install \"/kaggle/input/hpamisc/pycocotools-2.0-cp37-cp37m-linux_x86_64.whl\"\n#!pip install \"/kaggle/input/hpamisc/faiss_gpu-1.7.0-cp37-cp37m-manylinux2014_x86_64.whl\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:46:25.857648Z","iopub.execute_input":"2025-05-26T09:46:25.858058Z","iopub.status.idle":"2025-05-26T09:46:25.873418Z","shell.execute_reply.started":"2025-05-26T09:46:25.858015Z","shell.execute_reply":"2025-05-26T09:46:25.872418Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install \"../input/landmark-additional-packages/timm-0.4.12-py3-none-any.whl\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:46:25.87515Z","iopub.execute_input":"2025-05-26T09:46:25.87548Z","iopub.status.idle":"2025-05-26T09:47:04.766028Z","shell.execute_reply.started":"2025-05-26T09:46:25.875451Z","shell.execute_reply":"2025-05-26T09:47:04.764948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! python ../input/maozi-no-arcface/maozi_no_arcface.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:04.767912Z","iopub.execute_input":"2025-05-26T09:47:04.768334Z","iopub.status.idle":"2025-05-26T09:47:06.158408Z","shell.execute_reply.started":"2025-05-26T09:47:04.768284Z","shell.execute_reply":"2025-05-26T09:47:06.157311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append('../input/hpa-singlecell-e050f56/hpa_singlecell-double_level_valid_all/')\n\nfrom torch import nn\nimport torch\nimport torch.nn.functional as F\nimport torchvision\nimport timm\nfrom torch.nn.parameter import Parameter\nimport albumentations as A\n\nfrom utils import parse_args, prepare_for_result\nfrom torch.utils.data import DataLoader, Dataset\nfrom losses import get_loss, get_class_balanced_weighted\nfrom dataloaders import get_dataloader\nfrom utils import load_matched_state\nfrom configs import Config\nfrom models import get_model\nfrom dataloaders.transform_loader import get_tfms\n\ntensor_tfms = torchvision.transforms.Compose([\n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406, 0.406], std=[0.229, 0.224, 0.225, 0.225]),\n        ])\n\ntta_tfms = A.Compose([\n    A.Resize(always_apply=False, p=1, height=256, width=256, interpolation=1),\n    A.HorizontalFlip(always_apply=False, p=0.5),\n    A.ShiftScaleRotate(always_apply=False, p=0.7, shift_limit_x=(-0.06, 0.06), shift_limit_y=(-0.06, 0.06), scale_limit=(-0.3, 0.3), rotate_limit=(-22.5, 22.5), interpolation=1, border_mode=2, value=None, mask_value=None),\n    A.RandomBrightnessContrast(always_apply=False, p=0.5, brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), brightness_by_max=True),\n])\n\n\nimport base64\nimport zlib\nfrom pycocotools import _mask as coco_mask\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\nimport tqdm\nimport seaborn as sns","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:06.160347Z","iopub.execute_input":"2025-05-26T09:47:06.160785Z","iopub.status.idle":"2025-05-26T09:47:06.174228Z","shell.execute_reply.started":"2025-05-26T09:47:06.160744Z","shell.execute_reply":"2025-05-26T09:47:06.173007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def binary_mask_to_ascii(mask, mask_val=1):\n    \"\"\"Converts a binary mask into OID challenge encoding ascii text.\"\"\"\n    mask = np.where(mask==mask_val, 1, 0).astype(np.bool)\n    \n    # check input mask --\n    if mask.dtype != np.bool:\n        raise ValueError(f\"encode_binary_mask expects a binary mask, received dtype == {mask.dtype}\")\n\n    mask = np.squeeze(mask)\n    if len(mask.shape) != 2:\n        raise ValueError(f\"encode_binary_mask expects a 2d mask, received shape == {mask.shape}\")\n\n    # convert input mask to expected COCO API input --\n    mask_to_encode = mask.reshape(mask.shape[0], mask.shape[1], 1)\n    mask_to_encode = mask_to_encode.astype(np.uint8)\n    mask_to_encode = np.asfortranarray(mask_to_encode)\n\n    # RLE encode mask --\n    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n\n    # compress and base64 encoding --\n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    return base64_str.decode()\n\ndef process(x):\n    iid, msk, img, sz = x\n    img = cv2.resize(img, (2048, 2048))\n    enc_msk = cv2.resize(msk, (sz, sz))\n    cell_mask = msk\n    subs = {}\n    results = []\n    for i in range(1, cell_mask.max() + 1):\n        enc = binary_mask_to_ascii(enc_msk, i)\n        sub = cv2.resize((cell_mask == i).astype(np.float), (2048, 2048), cv2.INTER_LINEAR)\n        xr, yr = np.where(sub == 1)\n        xmin, xmax, ymin, ymax = xr.min(), xr.max(), yr.min(), yr.max()\n        subs[i] = (img * np.repeat((sub == 1).astype(np.int)[:, :, np.newaxis], 4, 2))[xmin:xmax, ymin: ymax]\n#         imsave(f'./seg_png_fix_test/{iid}_{i}.png', (255 * subs[i]).astype(np.uint8))\n        results.append(((255 * subs[i]).astype(np.uint8), enc, sz, sz))\n    return results\n\ndef squarify(M,val):\n    (a,b,c)=M.shape\n    if a>b:\n        padding=((0,0),((a-b)//2,a-b-(a-b)//2),(0, 0))\n    else:\n        padding=(((b-a)//2,b-a-(b-a)//2),(0,0),(0, 0))\n    return np.pad(M,padding,mode='constant',constant_values=val)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:06.175893Z","iopub.execute_input":"2025-05-26T09:47:06.176395Z","iopub.status.idle":"2025-05-26T09:47:06.200096Z","shell.execute_reply.started":"2025-05-26T09:47:06.176325Z","shell.execute_reply":"2025-05-26T09:47:06.199056Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# set up the train and rest data","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\n\n# 读取 CSV\ndf = pd.read_csv('/kaggle/input/hpa-single-cell-image-classification/train.csv')\n\n# 按 20% 训练, 80% 测试划分\ntrain_df, test_df = train_test_split(df, test_size=0.8, random_state=42)\n\n# 打印样本数\nprint(\"Train set:\", len(train_df), \"Test set:\", len(test_df))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:06.201602Z","iopub.execute_input":"2025-05-26T09:47:06.20201Z","iopub.status.idle":"2025-05-26T09:47:06.277392Z","shell.execute_reply.started":"2025-05-26T09:47:06.201969Z","shell.execute_reply":"2025-05-26T09:47:06.276379Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading models\n* b3\n* b5\n* r50d\n* r200d\n* se50","metadata":{}},{"cell_type":"code","source":"!ls ../input/hpa-single-cell-b3-philandrare-5f/5f_double_sin_exp5_rare.yaml/\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:06.27884Z","iopub.execute_input":"2025-05-26T09:47:06.279243Z","iopub.status.idle":"2025-05-26T09:47:07.534671Z","shell.execute_reply.started":"2025-05-26T09:47:06.279201Z","shell.execute_reply":"2025-05-26T09:47:07.533441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt = {\n    0: 13, 1: 12, 2: 12, 3: 11, 4: 14\n}\n\nmodels = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(5):\n    cfg = Config.load_json('../input/hpa-single-cell-b3-philandrare-5f/5f_double_sin_exp5_rare.yaml/config.json')\n    \n    model = get_model(cfg).to(device)\n    model_path = f'../input/hpa-single-cell-b3-philandrare-5f/5f_double_sin_exp5_rare.yaml/f{i}_epoch-{ckpt[i]}.pth'\n    load_matched_state(model, torch.load(model_path, map_location=device))\n    _ = model.eval()\n    models.append(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:07.536968Z","iopub.execute_input":"2025-05-26T09:47:07.537437Z","iopub.status.idle":"2025-05-26T09:47:10.659891Z","shell.execute_reply.started":"2025-05-26T09:47:07.537383Z","shell.execute_reply":"2025-05-26T09:47:10.658742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt = {\n    0: 18, 1: 14, 2: 14, 3: 15, 4: 15\n}\n\nmodels = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(5):\n    if i in [2, 3, 4]:\n        continue\n    cfg = Config.load_json('../input/hpa-b5-final-model/b5_final_hpa_0504/config.json')\n    model = get_model(cfg).to(device)\n    model_path = f'../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints/f{i}_epoch-{ckpt[i]}.pth'\n    load_matched_state(model, torch.load(model_path, map_location=device))\n    _ = model.eval()\n    models.append(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:10.661404Z","iopub.execute_input":"2025-05-26T09:47:10.661881Z","iopub.status.idle":"2025-05-26T09:47:13.17319Z","shell.execute_reply.started":"2025-05-26T09:47:10.661845Z","shell.execute_reply":"2025-05-26T09:47:13.172151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt = {\n    0: 19, 1: 19, 2: 17, 3: 17, 4: 18\n}\n\nmodels = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(5):\n    if i in [0, 3, 4]:\n        continue\n    cfg = Config.load_json('../input/hpa-resnet50d-0508/resnet50d_final/config.json')\n    model = get_model(cfg).to(device)\n    load_matched_state(model, torch.load(\n        f'../input/hpa-resnet50d-0508/resnet50d_final/checkpoints/f{i}_epoch-{ckpt[i]}.pth',\n        map_location=device))\n    _ = model.eval()\n    models.append(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:13.17481Z","iopub.execute_input":"2025-05-26T09:47:13.175115Z","iopub.status.idle":"2025-05-26T09:47:14.880181Z","shell.execute_reply.started":"2025-05-26T09:47:13.175084Z","shell.execute_reply":"2025-05-26T09:47:14.879242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp ../input/landmark-additional-packages/resnet200d_ra2-bdba9bf9.pth /root/.cache/torch/hub/checkpoints/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:14.881739Z","iopub.execute_input":"2025-05-26T09:47:14.882134Z","iopub.status.idle":"2025-05-26T09:47:17.204998Z","shell.execute_reply.started":"2025-05-26T09:47:14.882094Z","shell.execute_reply":"2025-05-26T09:47:17.203575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt = {\n    0: 15, 1: 15, 2: 13, 3: 13\n}\n\nmodels = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(4):\n    if i in [0, 1, 4]:  # 注意：你设定的是 range(4)，所以最大只能是 3，4 会被跳过\n        continue\n    cfg = Config.load_json('../input/hpa-jakiro-resnet200d/double_sin_exp5_r200d_rarex2_upload/config.json')\n    model = get_model(cfg).to(device)\n    load_matched_state(model, torch.load(\n        f'../input/hpa-jakiro-resnet200d/double_sin_exp5_r200d_rarex2_upload/f{i}_epoch-{ckpt[i]}.pth',\n        map_location=device))\n    _ = model.eval()\n    models.append(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:17.206861Z","iopub.execute_input":"2025-05-26T09:47:17.207245Z","iopub.status.idle":"2025-05-26T09:47:22.131768Z","shell.execute_reply.started":"2025-05-26T09:47:17.207201Z","shell.execute_reply":"2025-05-26T09:47:22.130713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt = {\n    0: 19, 1: 16, 2: 16, 3: 17, 4: 19\n}\n\nmodels = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(5):\n    if i in [0, 1, 2]: continue\n    print(i)\n    cfg = Config.load_json('../input/hpa-b5-final-model/b5_final_hpa_0504/config.json')\n    model = get_model(cfg).to(device)\n    load_matched_state(model, torch.load(\n        f'../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints/f{i}_epoch-{ckpt[i]}.pth',\n        map_location=device))\n    _ = model.eval()\n    models.append(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:22.13313Z","iopub.execute_input":"2025-05-26T09:47:22.133543Z","iopub.status.idle":"2025-05-26T09:47:24.564281Z","shell.execute_reply.started":"2025-05-26T09:47:22.133501Z","shell.execute_reply":"2025-05-26T09:47:24.563295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(models)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:24.565826Z","iopub.execute_input":"2025-05-26T09:47:24.56612Z","iopub.status.idle":"2025-05-26T09:47:24.571943Z","shell.execute_reply.started":"2025-05-26T09:47:24.566091Z","shell.execute_reply":"2025-05-26T09:47:24.570851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\n\n# Define best epochs for models 0 to 4\nckpt = {\n    0: 19, 1: 16, 2: 16, 3: 17, 4: 19\n}\n\nmodels = []\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nfor i in range(5):\n    model_path = f'../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints/f{i}_epoch-{ckpt[i]}.pth'\n    \n    if not os.path.exists(model_path):\n        print(f\"❌ Model {i} checkpoint not found: {model_path}\")\n        continue\n\n    print(f\"✅ Loading model {i} ...\")\n    cfg = Config.load_json('../input/hpa-b5-final-model/b5_final_hpa_0504/config.json')\n    model = get_model(cfg).to(device)\n    load_matched_state(model, torch.load(model_path, map_location=device))\n    model.eval()\n    models.append(model)\n\nprint(\"✅ Total models loaded:\", len(models))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:24.573149Z","iopub.execute_input":"2025-05-26T09:47:24.57343Z","iopub.status.idle":"2025-05-26T09:47:30.278957Z","shell.execute_reply.started":"2025-05-26T09:47:24.573401Z","shell.execute_reply":"2025-05-26T09:47:30.277765Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# load all models to see which one can fit -- i need a model to show the mask","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom PIL import Image\nimport os\nfrom torchvision.transforms.functional import normalize\nimport torch.nn.functional as F\n\n# 测试用图像 ID\ntest_image_id = '0040581b-f1f2-4fbe-b043-b6bfea5404bb'\nimage_dir = '/kaggle/input/hpa-single-cell-image-classification/test'\n\n# 加载 RGBy 图像\ndef load_rgby_image(image_id, image_dir, size=512):\n    def read_channel(color):\n        path = os.path.join(image_dir, f\"{image_id}_{color}.png\")\n        img = Image.open(path).convert(\"L\").resize((size, size))\n        return np.array(img)\n\n    red = read_channel(\"red\")\n    green = read_channel(\"green\")\n    blue = read_channel(\"blue\")\n    yellow = read_channel(\"yellow\")\n    rgby = np.stack([red, green, blue, yellow], axis=-1)\n    return rgby\n\n# 预处理 RGBy 图像：resize + normalize\ndef preprocess_rgby_np(np_img):\n    tensor = torch.from_numpy(np_img.transpose(2, 0, 1)).float() / 255.0  # [4, H, W]\n    tensor = F.interpolate(tensor.unsqueeze(0), size=(512, 512), mode='bilinear', align_corners=False).squeeze(0)\n    normed = normalize(tensor,\n                       mean=[0.485, 0.456, 0.406, 0.406],\n                       std=[0.229, 0.224, 0.225, 0.225])\n    return normed\n\n# 测试所有模型是否输出 [1, 1, H, W] 的 segmentation mask\nrgby = load_rgby_image(test_image_id, image_dir)\nimage_tensor = preprocess_rgby_np(rgby).unsqueeze(0).to(\"cpu\")\n\nvalid_model_indexes = []\n\nfor i, model in enumerate(models):\n    model = model.to(\"cpu\")\n    model.eval()\n    try:\n        with torch.no_grad():\n            output = model(image_tensor)\n            if output.ndim == 4 and output.shape[1] == 1:\n                valid_model_indexes.append(i)\n    except Exception as e:\n        print(f\"model[{i}] failed: {str(e)}\")\n\nprint(\"✅ segmentation-capable models:\", valid_model_indexes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:30.284518Z","iopub.execute_input":"2025-05-26T09:47:30.284842Z","iopub.status.idle":"2025-05-26T09:47:34.323067Z","shell.execute_reply.started":"2025-05-26T09:47:30.284812Z","shell.execute_reply":"2025-05-26T09:47:34.321976Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# load the model can do cell segmentation from mmdection","metadata":{}},{"cell_type":"code","source":"import base64, zlib\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport cv2\nimport pickle\nfrom pycocotools import mask as mutils\n\nROOT = '../input/hpa-single-cell-image-classification/'\ntrain_or_test = 'test'  # 或 'train' 视你的数据来源而定\n\n# 加载 RGB 图像（原始函数）\ndef read_img(image_id, color, train_or_test='train', image_size=None):\n    filename = f'{ROOT}/{train_or_test}/{image_id}_{color}.png'\n    assert os.path.exists(filename), f'not found {filename}'\n    img = cv2.imread(filename, cv2.IMREAD_UNCHANGED)\n    if image_size is not None:\n        img = cv2.resize(img, (image_size, image_size))\n    if img.dtype == 'uint16':\n        img = (img / 256).astype('uint8')\n    return img\n\ndef load_RGBY_image(image_id, train_or_test='train', image_size=None):\n    red = read_img(image_id, \"red\", train_or_test, image_size)\n    green = read_img(image_id, \"green\", train_or_test, image_size)\n    blue = read_img(image_id, \"blue\", train_or_test, image_size)\n    stacked_images = np.transpose(np.array([red, green, blue]), (1, 2, 0))\n    return stacked_images\n\ndef print_masked_img(image_id, mask):\n    img = load_RGBY_image(image_id, train_or_test)\n    plt.figure(figsize=(15, 5))\n    plt.subplot(1, 3, 1)\n    plt.imshow(img)\n    plt.title('Image')\n    plt.axis('off')\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(mask)\n    plt.title('Mask')\n    plt.axis('off')\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(img)\n    plt.imshow(mask, alpha=0.6)\n    plt.title('Image + Mask')\n    plt.axis('off')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:34.325154Z","iopub.execute_input":"2025-05-26T09:47:34.325465Z","iopub.status.idle":"2025-05-26T09:47:34.336419Z","shell.execute_reply.started":"2025-05-26T09:47:34.325434Z","shell.execute_reply":"2025-05-26T09:47:34.335271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/mmdetection-for-segmentation-inference-raw/mask_rcnn_resnest101_v5_ep9.pkl\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:34.337786Z","iopub.execute_input":"2025-05-26T09:47:34.338087Z","iopub.status.idle":"2025-05-26T09:47:35.591585Z","shell.execute_reply.started":"2025-05-26T09:47:34.338058Z","shell.execute_reply":"2025-05-26T09:47:35.590491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 正确加载结果\nresult = pickle.load(open('/kaggle/input/mmdetection-for-segmentation-inference-raw/mask_rcnn_resnest101_v5_ep9.pkl', 'rb'))\n\n# 你保存的图片 ID 和文件名列表（用于打印）\nannos = [{'filename': f\"{img_id}.jpg\"} for img_id in [\n    '004a2e56-b5e3-40a1-a579-113ed3f5769f',\n    '00865c52-7f2a-4fcb-8b89-b9b8c37c23d7'\n]]\n\n# 展示前三张图像的 mask\n# 展示前三张图像的 mask，并添加 image_id、class_id 和 confidence 到标题\nfor ii in range(3):\n    image_id = annos[ii]['filename'].replace('.jpg', '').replace('.png', '')\n    allowed_classes = b5_predicted_classes[image_id]\n\n    shown = 0  # 👈 加一个 counter 追踪展示次数\n\n    for class_id in allowed_classes:\n        bbs = result[ii][0][class_id]\n        sgs = result[ii][1][class_id]\n\n        for bb, sg in zip(bbs, sgs):\n            conf = bb[4]\n            if conf > 0.3:\n                print(f'class_id:{class_id}, image_id:{image_id}, confidence:{conf:.3f}')\n                mask = mutils.decode(sg).astype(bool)\n\n                img = load_RGBY_image(image_id, train_or_test)\n                plt.figure(figsize=(15, 5))\n\n                plt.subplot(1, 3, 1)\n                plt.imshow(img)\n                plt.title(f'Image\\n{image_id}')\n                plt.axis('off')\n\n                plt.subplot(1, 3, 2)\n                plt.imshow(mask)\n                plt.title(f'Mask\\nclass:{class_id}')\n                plt.axis('off')\n\n                plt.subplot(1, 3, 3)\n                plt.imshow(img)\n                plt.imshow(mask, alpha=0.6)\n                plt.title(f'Overlay\\nconf:{conf:.3f}')\n                plt.axis('off')\n\n                plt.tight_layout()\n                plt.show()\n\n                shown += 1\n                if shown == 3:  # 👈 达到展示数量就跳出\n                    break\n        if shown == 3:\n            break\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:24.381788Z","iopub.execute_input":"2025-05-26T09:49:24.382133Z","iopub.status.idle":"2025-05-26T09:49:24.413084Z","shell.execute_reply.started":"2025-05-26T09:49:24.382104Z","shell.execute_reply":"2025-05-26T09:49:24.411392Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# try to commect the b5 model with the mmdetection --personally predic it onlt has class 0 2 and 7","metadata":{}},{"cell_type":"code","source":"import pickle\n\nresult_path = '/kaggle/input/mmdetection-for-segmentation-inference-raw/mask_rcnn_resnest101_v5_ep9.pkl'\nwith open(result_path, 'rb') as f:\n    result = pickle.load(f)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:32.468509Z","iopub.execute_input":"2025-05-26T09:49:32.468939Z","iopub.status.idle":"2025-05-26T09:49:32.477449Z","shell.execute_reply.started":"2025-05-26T09:49:32.468904Z","shell.execute_reply":"2025-05-26T09:49:32.47648Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# choose one sample see if can have the result","metadata":{}},{"cell_type":"code","source":"# 示例：选取一个样本，看看模型是否能输出预测\nsample_id = train_df.iloc[0]['ID']\nprint(\"Trying image:\", sample_id)\n\n# 构造路径，读取红绿蓝通道图像（按照当前 notebook 的规范）\nred_path = f'/kaggle/input/hpa-single-cell-image-classification/train/{sample_id}_red.png'\ngreen_path = f'/kaggle/input/hpa-single-cell-image-classification/train/{sample_id}_green.png'\nblue_path = f'/kaggle/input/hpa-single-cell-image-classification/train/{sample_id}_blue.png'\nyellow_path = f'/kaggle/input/hpa-single-cell-image-classification/train/{sample_id}_yellow.png'\n\n# 用 cv2 加载图像\nimport cv2\nimport numpy as np\n\nr = cv2.imread(red_path, 0).astype(np.float32) / 255.\ng = cv2.imread(green_path, 0).astype(np.float32) / 255.\nb = cv2.imread(blue_path, 0).astype(np.float32) / 255.\ny = cv2.imread(yellow_path, 0).astype(np.float32) / 255.\n\nimg = np.stack([r, g, b, y], axis=-1)  # (H, W, 4)\nprint(\"Image shape:\", img.shape)\n\n# 如果有 transforms 或模型输入结构，需 resize + normalize\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:35.548671Z","iopub.execute_input":"2025-05-26T09:49:35.549011Z","iopub.status.idle":"2025-05-26T09:49:35.971104Z","shell.execute_reply.started":"2025-05-26T09:49:35.548982Z","shell.execute_reply":"2025-05-26T09:49:35.9701Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# predict result","metadata":{}},{"cell_type":"code","source":"import os\nos.listdir('/kaggle/input/hpa-single-cell-image-classification/train')[:10]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:40.550881Z","iopub.execute_input":"2025-05-26T09:49:40.551244Z","iopub.status.idle":"2025-05-26T09:49:40.600214Z","shell.execute_reply.started":"2025-05-26T09:49:40.551213Z","shell.execute_reply":"2025-05-26T09:49:40.599107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfilenames = os.listdir('/kaggle/input/hpa-single-cell-image-classification/train')\nids = set(fname.split('_')[0] for fname in filenames)\nprint(len(ids), \"unique image IDs found\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:44.508352Z","iopub.execute_input":"2025-05-26T09:49:44.50871Z","iopub.status.idle":"2025-05-26T09:49:44.596075Z","shell.execute_reply.started":"2025-05-26T09:49:44.508682Z","shell.execute_reply":"2025-05-26T09:49:44.594951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 导入必要模块\nimport os\nimport cv2\nimport numpy as np\nimport torch\n\n# 获取训练集中所有图片 ID\ntrain_folder = '/kaggle/input/hpa-single-cell-image-classification/train'\nfilenames = os.listdir(train_folder)\nids = sorted(set(fname.split('_')[0] for fname in filenames))\n\n# 选取第一个 sample id\nsample_id = ids[0]\nprint(\"Using sample:\", sample_id)\n\n# 构造图像路径\ndef build_path(color):\n    return os.path.join(train_folder, f\"{sample_id}_{color}.png\")\n\nred_path = build_path('red')\ngreen_path = build_path('green')\nblue_path = build_path('blue')\nyellow_path = build_path('yellow')\n\n# 读取图像并缩放为 512x512（避免模型尺寸不匹配）\ndef load_and_resize(path):\n    img = cv2.imread(path, 0)  # 灰度读取\n    img = cv2.resize(img, (512, 512))\n    return img.astype(np.float32) / 255.\n\nr = load_and_resize(red_path)\ng = load_and_resize(green_path)\nb = load_and_resize(blue_path)\ny = load_and_resize(yellow_path)\n\n# 构建输入张量（B, C, H, W）\nimg = np.stack([r, g, b, y], axis=-1)  # (H, W, 4)\nimg = np.transpose(img, (2, 0, 1))  # (4, H, W)\nimg_tensor = torch.tensor(img).unsqueeze(0).float().to(device)\n\n# 使用模型进行预测\nwith torch.no_grad():\n    pred = models[0](img_tensor, cnt=1)[0]\nprint(\"Prediction shape:\", pred.shape)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:47.644684Z","iopub.execute_input":"2025-05-26T09:49:47.645048Z","iopub.status.idle":"2025-05-26T09:49:48.645719Z","shell.execute_reply.started":"2025-05-26T09:49:47.645019Z","shell.execute_reply":"2025-05-26T09:49:48.644494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predicted_class = pred.argmax(dim=1).item()\nprint(\"Predicted class:\", predicted_class)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:53.982705Z","iopub.execute_input":"2025-05-26T09:49:53.983109Z","iopub.status.idle":"2025-05-26T09:49:53.989306Z","shell.execute_reply.started":"2025-05-26T09:49:53.983073Z","shell.execute_reply":"2025-05-26T09:49:53.98835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"probs = torch.softmax(pred, dim=1)\ntop_probs, top_classes = torch.topk(probs, 5)\nprint(\"Top 5 class indices:\", top_classes)\nprint(\"Top 5 probabilities:\", top_probs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:56.587626Z","iopub.execute_input":"2025-05-26T09:49:56.587968Z","iopub.status.idle":"2025-05-26T09:49:56.596556Z","shell.execute_reply.started":"2025-05-26T09:49:56.587938Z","shell.execute_reply":"2025-05-26T09:49:56.59547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\n\n# ----------------------------------------\n# 函数1：加载RGB图像（从RGBY图中提取前三通道）\n# ----------------------------------------\ndef load_RGB_image(image_id, \n                   root='../input/hpa-single-cell-image-classification', \n                   train_or_test='test', \n                   size=None):\n    def read_channel(color):\n        path = os.path.join(root, train_or_test, f\"{image_id}_{color}.png\")\n        img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n        if size:\n            img = cv2.resize(img, (size, size))\n        return img\n\n    red = read_channel('red')\n    green = read_channel('green')\n    blue = read_channel('blue')\n    rgb = np.stack([red, green, blue], axis=-1)\n    return rgb\n\n# ----------------------------------------\n# 函数2：展示预测结果（图像 + 掩码 + 叠加）\n# ----------------------------------------\ndef show_prediction(image, mask, class_id, image_id, confidence):\n    plt.figure(figsize=(12, 4))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(image)\n    plt.title(\"Image\")\n    plt.axis('off')\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(mask.astype(np.uint8), cmap='gray')  # 或使用 cmap='Reds'\n    plt.title(f\"Mask\\nclass:{class_id}\")\n    plt.axis('off')\n\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(image)\n    plt.imshow(mask, alpha=0.6)\n    plt.title(f\"Overlay\\nclass:{class_id}, ID:{image_id}\\nConf:{confidence:.3f}\")\n    plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# ----------------------------------------\n# 示例数据（用于测试，后续用模型输出替换）\n# ----------------------------------------\nexample_predictions = [\n    {\n        'image_id': '0040581b-f1f2-4fbe-b043-b6bfea5404bb',\n        'mask': np.zeros((512, 512)),  # ← 替换成真实模型的mask输出\n        'class_id': 0,\n        'confidence': 0.47\n    },\n    {\n        'image_id': '0040581b-f1f2-4fbe-b043-b6bfea5404bb',\n        'mask': np.zeros((512, 512)),  # ← 第二个预测\n        'class_id': 0,\n        'confidence': 0.46\n    },\n]\n\n# ----------------------------------------\n# 主执行：循环展示所有结果\n# ----------------------------------------\nfor pred in example_predictions:\n    img = load_RGB_image(pred['image_id'])\n    show_prediction(img, pred['mask'], pred['class_id'], pred['image_id'], pred['confidence'])\n\nprint(\"Mask sum:\", mask.sum())\nprint(f\"Mask dtype: {mask.dtype}, shape: {mask.shape}, unique: {np.unique(mask)}\")\nprint(f\"Mask sum: {mask.sum()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:49:59.892533Z","iopub.execute_input":"2025-05-26T09:49:59.89289Z","iopub.status.idle":"2025-05-26T09:50:02.522971Z","shell.execute_reply.started":"2025-05-26T09:49:59.892861Z","shell.execute_reply":"2025-05-26T09:50:02.521914Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ask the b5 modele to do the predictin first","metadata":{}},{"cell_type":"code","source":"!ls ../input/hpa-b5-final-model/b5_final_hpa_0504\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:50:10.794518Z","iopub.execute_input":"2025-05-26T09:50:10.794889Z","iopub.status.idle":"2025-05-26T09:50:12.039689Z","shell.execute_reply.started":"2025-05-26T09:50:10.794858Z","shell.execute_reply":"2025-05-26T09:50:12.03855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls ../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:50:14.15694Z","iopub.execute_input":"2025-05-26T09:50:14.157317Z","iopub.status.idle":"2025-05-26T09:50:15.409827Z","shell.execute_reply.started":"2025-05-26T09:50:14.157279Z","shell.execute_reply":"2025-05-26T09:50:15.408521Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# we can conclue that the f3_epoch-19 has a better result","metadata":{}},{"cell_type":"code","source":"config_path = '../input/hpa-b5-final-model/b5_final_hpa_0504/config.json'\nmodel_path = '../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints/f3_epoch-19.pth'\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:50:20.002423Z","iopub.execute_input":"2025-05-26T09:50:20.002818Z","iopub.status.idle":"2025-05-26T09:50:20.007491Z","shell.execute_reply.started":"2025-05-26T09:50:20.002778Z","shell.execute_reply":"2025-05-26T09:50:20.00638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ✅ 图像预处理设置\nimport os\nimport cv2\nimport numpy as np\nfrom torchvision import transforms\n\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: torch.from_numpy(x.transpose((2, 0, 1))).float() / 255.0),\n    transforms.Resize((512, 512)),\n    transforms.Normalize([0.5]*4, [0.25]*4)  # 假设 RGBY 每通道都做简单归一化\n])\n\n\n# ✅ 定义函数：模型预测图像的标签\ndef get_b5_predicted_classes(model, image_ids, img_dir, threshold=0.5):\n    results = {}\n    for image_id in image_ids:\n        def read_channel(channel):\n            path = os.path.join(img_dir, f\"{image_id}_{channel}.png\")\n            img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n            return cv2.resize(img, (512, 512))\n\n        r = read_channel('red')\n        g = read_channel('green')\n        b = read_channel('blue')\n        y = read_channel('yellow')  # 新增\n        \n        # 合并 4 通道：RGBY\n        rgby = np.stack([r, g, b, y], axis=-1).astype(np.uint8)\n        rgby_tensor = transform(rgby).unsqueeze(0)\n    \n    with torch.no_grad():\n        logits, _ = model(rgby_tensor, cnt=1)\n        probs = torch.sigmoid(logits).squeeze().numpy()\n        print(f\"{image_id} probs:\", probs.round(3))\n        \n        topk = 3\n        pred_classes = probs.argsort()[-topk:][::-1].tolist()\n        print(f\"{image_id} top-{topk} predicted classes:\", pred_classes)\n        \n        results[image_id] = pred_classes\n    return results\n\n# ✅ 使用模型对图像进行分类预测\nimg_dir = '/kaggle/input/hpa-single-cell-image-classification/test'\nimage_ids = sorted(set(f.split('_')[0] for f in os.listdir(img_dir)))[:18]\n\n\n\nb5_predicted_classes = get_b5_predicted_classes(model, image_ids, img_dir, threshold=0.2)\n\n# ✅ 查看预测结果\nprint(b5_predicted_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:50:22.75817Z","iopub.execute_input":"2025-05-26T09:50:22.758607Z","iopub.status.idle":"2025-05-26T09:50:28.90055Z","shell.execute_reply.started":"2025-05-26T09:50:22.758568Z","shell.execute_reply":"2025-05-26T09:50:28.899412Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# if i try average about b 5 model","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom torchvision import transforms\n\n# ✅ 图像预处理方式\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: torch.from_numpy(x.transpose((2, 0, 1))).float() / 255.0),\n    transforms.Resize((512, 512)),\n    transforms.Normalize([0.5]*4, [0.25]*4)\n])\n\n# ✅ 手动设置\nconfig_path = '../input/hpa-b5-final-model/b5_final_hpa_0504/config.json'\ncheckpoint_dir = '../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints'\nimg_dir = '/kaggle/input/hpa-single-cell-image-classification/test'\nimage_ids = sorted(set(f.split('_')[0] for f in os.listdir(img_dir)))[:4]\n\n# ✅ 加载模型函数（复制你原来成功的版本）\ncfg_template = Config.load_json(config_path)\n\ndef load_model(model_path, device='cpu'):\n    model = get_model(cfg_template).to(device)\n    state_dict = torch.load(model_path, map_location=device)\n    model.load_state_dict(state_dict)\n    model.eval()\n    return model\n\n# ✅ 寻找存在的 checkpoint（优先用 epoch-19，如果没有往下找）\ndef find_existing_checkpoint(fold, max_epoch=19):\n    for epoch in reversed(range(max_epoch + 1)):\n        candidate = f\"f{fold}_epoch-{epoch}.pth\"\n        full_path = os.path.join(checkpoint_dir, candidate)\n        if os.path.exists(full_path):\n            return full_path\n    raise FileNotFoundError(f\"No checkpoint found for fold {fold}\")\n\n# ✅ 多模型平均预测\ndef ensemble_predict(image_ids, folds=[0,1,2,3,4], cnt=1, topk=3):\n    results = {}\n    for image_id in image_ids:\n        def read_channel(channel):\n            path = os.path.join(img_dir, f\"{image_id}_{channel}.png\")\n            img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n            return cv2.resize(img, (512, 512))\n\n        r = read_channel('red')\n        g = read_channel('green')\n        b = read_channel('blue')\n        y = read_channel('yellow')\n        rgby = np.stack([r, g, b, y], axis=-1).astype(np.uint8)\n        rgby_tensor = transform(rgby).unsqueeze(0)\n\n        logits_list = []\n\n        for fold in folds:\n            model_path = find_existing_checkpoint(fold)\n            print(f\"✔️ Using model: {model_path}\")\n            model = load_model(model_path)\n\n            with torch.no_grad():\n                logits, _ = model(rgby_tensor, cnt=cnt)\n                logits_list.append(torch.sigmoid(logits).squeeze().numpy())\n\n        avg_probs = np.mean(logits_list, axis=0)\n        print(f\"{image_id} avg_probs:\", avg_probs.round(3))\n\n        pred_classes = avg_probs.argsort()[-topk:][::-1].tolist()\n        print(f\"{image_id} top-{topk} ensemble predicted:\", pred_classes)\n        results[image_id] = pred_classes\n    return results\n\n# ✅ 执行预测\nb5_ensemble_predicted_classes = ensemble_predict(image_ids)\nprint(\"\\n✅ Final Ensemble Prediction:\\n\", b5_ensemble_predicted_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:51:38.629181Z","iopub.execute_input":"2025-05-26T09:51:38.629561Z","iopub.status.idle":"2025-05-26T09:52:14.795482Z","shell.execute_reply.started":"2025-05-26T09:51:38.629529Z","shell.execute_reply":"2025-05-26T09:52:14.794211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport cv2\nfrom torchvision import transforms\n\n# ✅ 图像预处理方式\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: torch.from_numpy(x.transpose((2, 0, 1))).float() / 255.0),\n    transforms.Resize((512, 512)),\n    transforms.Normalize([0.5]*4, [0.25]*4)\n])\n\n# ✅ 手动设置\nconfig_path = '../input/hpa-b5-final-model/b5_final_hpa_0504/config.json'\ncheckpoint_dir = '../input/hpa-b5-final-model/b5_final_hpa_0504/checkpoints'\nimg_dir = '/kaggle/input/hpa-single-cell-image-classification/test'\nimage_ids = sorted(set(f.split('_')[0] for f in os.listdir(img_dir)))[:4]\n\n# ✅ 加载模型函数\ncfg_template = Config.load_json(config_path)\n\ndef load_model(model_path, device='cpu'):\n    model = get_model(cfg_template).to(device)\n    state_dict = torch.load(model_path, map_location=device)\n    model.load_state_dict(state_dict)\n    model.eval()\n    return model\n\n# ✅ 寻找可用 checkpoint（默认 epoch 从 19 往前找）\ndef find_existing_checkpoint(fold, max_epoch=19):\n    for epoch in reversed(range(max_epoch + 1)):\n        candidate = f\"f{fold}_epoch-{epoch}.pth\"\n        full_path = os.path.join(checkpoint_dir, candidate)\n        if os.path.exists(full_path):\n            return full_path\n    raise FileNotFoundError(f\"No checkpoint found for fold {fold}\")\n\n# ✅ 多模型平均 + 混合预测（topk + threshold）\ndef ensemble_predict(image_ids, folds=[0,1,2,3,4], cnt=1, topk=3, threshold=0.1):\n    results = {}\n    for image_id in image_ids:\n        def read_channel(channel):\n            path = os.path.join(img_dir, f\"{image_id}_{channel}.png\")\n            img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n            return cv2.resize(img, (512, 512))\n\n        r = read_channel('red')\n        g = read_channel('green')\n        b = read_channel('blue')\n        y = read_channel('yellow')\n        rgby = np.stack([r, g, b, y], axis=-1).astype(np.uint8)\n        rgby_tensor = transform(rgby).unsqueeze(0)\n\n        logits_list = []\n\n        for fold in folds:\n            model_path = find_existing_checkpoint(fold)\n            print(f\"✔️ Using model: {model_path}\")\n            model = load_model(model_path)\n\n            with torch.no_grad():\n                logits, _ = model(rgby_tensor, cnt=cnt)\n                logits_list.append(torch.sigmoid(logits).squeeze().numpy())\n\n        avg_probs = np.mean(logits_list, axis=0)\n        print(f\"{image_id} avg_probs:\", avg_probs.round(3))\n\n        # ✅ 混合策略：topk + threshold\n        top_classes = avg_probs.argsort()[-topk:][::-1].tolist()\n        above_threshold = np.where(avg_probs > threshold)[0].tolist()\n        pred_classes = sorted(set(top_classes + above_threshold))\n\n        print(f\"{image_id} final predicted classes:\", pred_classes)\n        results[image_id] = pred_classes\n    return results\n\n# ✅ 执行预测\nb5_ensemble_predicted_classes = ensemble_predict(image_ids)\nprint(\"\\n✅ Final Ensemble Prediction:\\n\", b5_ensemble_predicted_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:52:22.949026Z","iopub.execute_input":"2025-05-26T09:52:22.949424Z","iopub.status.idle":"2025-05-26T09:52:58.490318Z","shell.execute_reply.started":"2025-05-26T09:52:22.949354Z","shell.execute_reply":"2025-05-26T09:52:58.489238Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# from a mask file","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport os\nimport cv2\nimport pickle\nfrom tqdm import tqdm\nfrom pycocotools import mask as mutils\n\n# ✅ 参数设置\nresult = pickle.load(open('/kaggle/input/mmdetection-for-segmentation-inference-raw/mask_rcnn_resnest101_v5_ep9.pkl', 'rb'))\n\n# 使用你之前跑过 B5 模型后的预测结果\nimage_ids = list(b5_ensemble_predicted_classes.keys())\n\n# 你之前加载 annos 的方式\nannos = [{'filename': f\"{img_id}.jpg\"} for img_id in image_ids]\n\n# ✅ 建立 image_id → result index 的映射（防止 index 错误）\nimageid_to_idx = {\n    anno['filename'].replace('.jpg', '').replace('.png', ''): i\n    for i, anno in enumerate(annos)\n}\n\n# ✅ 输出保存路径\noutput_dir = \"./filtered_masks_npz\"\nos.makedirs(output_dir, exist_ok=True)\n\n# ✅ 生成 instance 掩膜 + label dict，并保存为 .npz\nfor image_id in tqdm(image_ids):\n    if image_id not in imageid_to_idx:\n        print(f\"❌ Skipping {image_id}: not found in result.\")\n        continue\n\n    idx = imageid_to_idx[image_id]\n    allowed_classes = b5_ensemble_predicted_classes[image_id]\n    height, width = 512, 512\n\n    instance_mask = np.zeros((height, width), dtype=np.uint16)\n    label_dict = {}\n    instance_id = 1\n\n    for class_id in allowed_classes:\n        try:\n            bbs = result[idx][0][class_id]\n            sgs = result[idx][1][class_id]\n        except IndexError:\n            continue\n\n        for bb, sg in zip(bbs, sgs):\n            conf = bb[4]\n            if conf > 0.3:\n                mask = mutils.decode(sg).astype(bool)\n                mask = cv2.resize(mask.astype(np.uint8), (width, height), interpolation=cv2.INTER_NEAREST).astype(bool)\n                instance_mask[mask] = instance_id\n                label_dict[instance_id] = class_id\n                instance_id += 1\n\n    save_path = os.path.join(output_dir, f\"{image_id}.npz\")\n    np.savez_compressed(save_path, instance_mask, label_dict)\n    print(f\"✅ Saved {image_id}.npz with {instance_id-1} instances\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:04:35.937154Z","iopub.execute_input":"2025-05-26T10:04:35.937527Z","iopub.status.idle":"2025-05-26T10:04:36.893434Z","shell.execute_reply.started":"2025-05-26T10:04:35.937497Z","shell.execute_reply":"2025-05-26T10:04:36.892397Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# connect with the mask","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport numpy as np\nimport cv2\nimport os\n\n# ✅ cell type 映射（你可根据需要替换中文或完整标签）\nLABEL_MAP = {\n    0: 'Nucleoplasm',\n    1: 'Nuclear membrane',\n    2: 'Nucleoli',\n    3: 'Nucleoli fibrillar center',\n    4: 'Nuclear speckles',\n    5: 'Nuclear bodies',\n    6: 'Endoplasmic reticulum',\n    7: 'Golgi apparatus',\n    8: 'Peroxisomes',\n    9: 'Endosomes',\n    10: 'Lysosomes',\n    11: 'Intermediate filaments',\n    12: 'Actin filaments',\n    13: 'Focal adhesion sites',\n    14: 'Microtubules',\n    15: 'Microtubule ends',\n    16: 'Cytosol',\n    17: 'Plasma membrane',\n    18: 'Mitochondria',\n}\n\n# ✅ 可视化函数：一行三图（原图 / 筛选后掩膜 / 叠加图）\ndef visualize_mask_by_pred(image_id, predicted_classes, img_dir, mask_dir):\n    def read_channel(image_id, channel):\n        path = os.path.join(img_dir, f\"{image_id}_{channel}.png\")\n        img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n        return cv2.resize(img, (512, 512))\n    \n    r = read_channel(image_id, 'red')\n    g = read_channel(image_id, 'green')\n    b = read_channel(image_id, 'blue')\n    rgb = np.stack([r, g, b], axis=-1).astype(np.uint8)\n\n    # ✅ 加载掩膜（假设保存为 .npz 文件，包含 mask 和对应 label）\n    mask_path = os.path.join(mask_dir, f\"{image_id}.npz\")\n    data = np.load(mask_path, allow_pickle=True)\n    mask = data['arr_0']          # 每个像素的 instance id\n    label_map = data['arr_1'].item()  # dict: instance_id → class_id\n\n    # ✅ 构造筛选后的 binary mask\n    h, w = mask.shape\n    filtered_mask = np.zeros((h, w), dtype=np.uint8)\n    for instance_id, class_id in label_map.items():\n        if class_id in predicted_classes:\n            filtered_mask[mask == instance_id] = class_id + 1  # 显示用 class_id+1\n\n    # ✅ 可视化\n    plt.figure(figsize=(15, 5))\n    \n    plt.subplot(1, 3, 1)\n    plt.imshow(rgb)\n    plt.title(\"Original Image\")\n    plt.axis(False)\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(filtered_mask, cmap='tab20')\n    plt.title(f\"Filtered Mask\\n{[LABEL_MAP[i] for i in predicted_classes]}\")\n    plt.axis(False)\n\n    plt.subplot(1, 3, 3)\n    overlay = rgb.copy()\n    overlay[filtered_mask > 0] = (255, 0, 0)  # 红色覆盖\n    plt.imshow(overlay)\n    plt.title(\"Overlay Mask\")\n    plt.axis(False)\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:05:53.084168Z","iopub.execute_input":"2025-05-26T10:05:53.08474Z","iopub.status.idle":"2025-05-26T10:05:53.100518Z","shell.execute_reply.started":"2025-05-26T10:05:53.084698Z","shell.execute_reply":"2025-05-26T10:05:53.099152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 可视化输出路径\nmask_dir = \"./filtered_masks_npz\"\n\n# 显示其中一张图像\nimage_id = '0040581b-f1f2-4fbe-b043-b6bfea5404bb'  # 你可以换成其他 ID\npred_classes = b5_ensemble_predicted_classes[image_id]\n\nvisualize_mask_by_pred(image_id, pred_classes, img_dir, mask_dir)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:06:18.191245Z","iopub.execute_input":"2025-05-26T10:06:18.191626Z","iopub.status.idle":"2025-05-26T10:06:18.849203Z","shell.execute_reply.started":"2025-05-26T10:06:18.191594Z","shell.execute_reply":"2025-05-26T10:06:18.848256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def bootstrap_ci(data, n_boot=1000, ci=0.95):\n    samples = np.random.choice(data, (n_boot, len(data)), replace=True)\n    means = samples.mean(axis=1)\n    lower = np.percentile(means, (1 - ci) / 2 * 100)\n    upper = np.percentile(means, (1 + ci) / 2 * 100)\n    return np.mean(data), lower, upper\n\ndef ensemble_predict_with_ci(image_ids, checkpoint_dir, cfg_template, folds=[0,1,2,3,4], cnt=1, topk=3, threshold=0.1):\n    class_prob_collection = {i: [] for i in range(19)}\n    results = {}\n\n    for image_id in image_ids:\n        # Load image\n        def read_channel(channel):\n            path = os.path.join(img_dir, f\"{image_id}_{channel}.png\")\n            img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n            return cv2.resize(img, (512, 512))\n\n        r = read_channel('red')\n        g = read_channel('green')\n        b = read_channel('blue')\n        y = read_channel('yellow')\n        rgby = np.stack([r, g, b, y], axis=-1).astype(np.uint8)\n        rgby_tensor = transform(rgby).unsqueeze(0)\n\n        # Ensemble prediction\n        logits_list = []\n        for fold in folds:\n            model_path = find_existing_checkpoint(fold)\n            model = load_model(model_path)\n            with torch.no_grad():\n                logits, _ = model(rgby_tensor, cnt=cnt)\n                logits_list.append(torch.sigmoid(logits).squeeze().numpy())\n\n        avg_probs = np.mean(logits_list, axis=0)\n        for i in range(19):\n            class_prob_collection[i].append(avg_probs[i])\n\n        # Threshold + topk strategy\n        top_classes = avg_probs.argsort()[-topk:][::-1].tolist()\n        above_threshold = np.where(avg_probs > threshold)[0].tolist()\n        pred_classes = sorted(set(top_classes + above_threshold))\n        results[image_id] = pred_classes\n\n    # Compute CI\n    ci_dict = {i: bootstrap_ci(class_prob_collection[i]) for i in range(19)}\n    return results, ci_dict\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:40:07.557781Z","iopub.execute_input":"2025-05-26T10:40:07.558155Z","iopub.status.idle":"2025-05-26T10:40:07.570939Z","shell.execute_reply.started":"2025-05-26T10:40:07.558126Z","shell.execute_reply":"2025-05-26T10:40:07.569945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"b5_ensemble_predicted_classes, ci_dict = ensemble_predict_with_ci(\n    image_ids=image_ids,\n    checkpoint_dir=checkpoint_dir,\n    cfg_template=cfg_template,\n    folds=[0, 1, 2, 3, 4]\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:38:30.130917Z","iopub.execute_input":"2025-05-26T10:38:30.13126Z","iopub.status.idle":"2025-05-26T10:39:05.346772Z","shell.execute_reply.started":"2025-05-26T10:38:30.131231Z","shell.execute_reply":"2025-05-26T10:39:05.345594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch\n\nLABEL_MAP = {\n    0: 'Nucleoplasm', 1: 'Nuclear membrane', 2: 'Nucleoli',\n    3: 'Nucleoli fibrillar center', 4: 'Nuclear speckles', 5: 'Nuclear bodies',\n    6: 'Endoplasmic reticulum', 7: 'Golgi apparatus', 8: 'Peroxisomes',\n    9: 'Endosomes', 10: 'Lysosomes', 11: 'Intermediate filaments',\n    12: 'Actin filaments', 13: 'Focal adhesion sites', 14: 'Microtubules',\n    15: 'Microtubule ends', 16: 'Cytosol', 17: 'Plasma membrane', 18: 'Mitochondria'\n}\n\ndef show_ci_legend(image_id):\n    pred_classes = b5_ensemble_predicted_classes[image_id]\n    legend_elements = []\n    for class_id in pred_classes:\n        label = LABEL_MAP[class_id]\n        mean, low, high = ci_dict[class_id]\n        ci_text = f\"{mean:.2f} (95% CI: {low:.2f}-{high:.2f})\"\n        legend_elements.append(Patch(facecolor='red', edgecolor='black', label=f\"{label}: {ci_text}\"))\n    plt.figure(figsize=(8, 2))\n    plt.legend(handles=legend_elements, loc='center', frameon=True)\n    plt.axis('off')\n    plt.title(f\"Predicted Cell Types with CI for {image_id}\")\n    plt.tight_layout()\n    plt.show()\n\n# 示例调用\nshow_ci_legend('0040581b-f1f2-4fbe-b043-b6bfea5404bb')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T10:40:14.804542Z","iopub.execute_input":"2025-05-26T10:40:14.804887Z","iopub.status.idle":"2025-05-26T10:40:14.970116Z","shell.execute_reply.started":"2025-05-26T10:40:14.804858Z","shell.execute_reply":"2025-05-26T10:40:14.968746Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## If we read from a csv","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('submission.csv')\n\nimgs = []\nfor i, x in df.iterrows():\n    label = x.PredictionString.split(' ')[0::3]\n    prob = x.PredictionString.split(' ')[1::3]\n    encodes = x.PredictionString.split(' ')[2::3]\n    for idx, enc in enumerate(list(set(encodes))):\n        imgs.append({\n            'image_id': x.ID,\n            'cell_id': idx+1,\n            'enc': enc,\n            'fname': f'{x.ID}_{idx+1}',\n        })\n\ntm = pd.DataFrame(imgs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.667473Z","iopub.status.idle":"2025-05-26T09:47:35.667951Z","shell.execute_reply":"2025-05-26T09:47:35.667724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"probs = []\nfor i, x in df.iterrows():\n    label = x.PredictionString.split(' ')[0::3]\n    prob = x.PredictionString.split(' ')[1::3]\n    encodes = x.PredictionString.split(' ')[2::3]\n    for idx, enc in enumerate(encodes):\n        probs.append({\n            'enc': enc,\n            'predict': int(label[idx]),\n            'prob': float(prob[idx])\n        })\n\nprob = pd.DataFrame(probs)\ntm_pred = prob.groupby(['enc', 'predict']).mean().unstack()['prob']\ntm_pred.columns.name = ''\nteam = tm[['enc', 'fname']].merge(tm_pred.reset_index(), on='enc', how='inner').drop('enc', 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.668837Z","iopub.status.idle":"2025-05-26T09:47:35.669252Z","shell.execute_reply":"2025-05-26T09:47:35.669037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission = pd.read_csv('../input/hpa-single-cell-image-classification/sample_submission.csv', index_col=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.670108Z","iopub.status.idle":"2025-05-26T09:47:35.670535Z","shell.execute_reply":"2025-05-26T09:47:35.670305Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"team_pred = team.set_index('fname')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.671464Z","iopub.status.idle":"2025-05-26T09:47:35.67187Z","shell.execute_reply":"2025-05-26T09:47:35.671674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SliceInferenceDataset(torch.utils.data.Dataset):\n    def __init__(self, df, tta=16, cfg=None, tfms=None):\n        self.df = df\n        self.iids = self.df.image_id.unique()\n        self.tta = tta\n        \n    def __len__(self):\n        return len(self.iids)\n\n    def __getitem__(self, idx):\n        iid = self.iids[idx]\n        mt = f'../input/hpa-single-cell-image-classification/test/{iid}_red.png'\n        er = f'../input/hpa-single-cell-image-classification/test/{iid}_yellow.png'\n        nu = f'../input/hpa-single-cell-image-classification/test/{iid}_blue.png'\n        pr = f'../input/hpa-single-cell-image-classification/test/{iid}_green.png'\n        r = cv2.imread(mt, 0).astype(np.float) / 255.0\n        g = cv2.imread(pr, 0).astype(np.float) / 255.0\n        b = cv2.imread(nu, 0).astype(np.float) / 255.0\n        a = cv2.imread(er, 0).astype(np.float) / 255.0\n        sz = r.shape[0]\n        img = np.stack([r, g, b, a], -1)\n        sli = []\n        for i, x in self.df[self.df.image_id == iid].iterrows():\n            bd = base64.b64decode(x.enc)\n            zd = zlib.decompress(bd)\n            encoded = [{'counts': zd, 'size': (sz, sz)}]\n            ded = coco_mask.decode(encoded)[:, :, 0]\n\n            xr, yr = np.where(ded == 1)\n            sub = img[xr.min(): xr.max(), yr.min(): yr.max()]\n            crop_sub_mask = ded[xr.min(): xr.max(), yr.min(): yr.max()]\n            crop_sub_mask = np.repeat(crop_sub_mask[:, :, np.newaxis], 4, axis=2)\n            r = sub * crop_sub_mask\n            sli.append((cv2.resize(squarify(r, 0), (256, 256)).astype(np.float32), x.fname))\n        BS, tta=len(sli) + 1, self.tta\n        ipts = []\n        raw_ipt = [e[0] for e in sli]\n        for tt in range(tta):\n            ipts.append(torch.stack([tensor_tfms(tta_tfms(image=x)['image']) for x in raw_ipt]).float())\n        return ipts, BS, len(sli), tta, iid, [x[1] for x in sli]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.672742Z","iopub.status.idle":"2025-05-26T09:47:35.673132Z","shell.execute_reply":"2025-05-26T09:47:35.672936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# tm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.674121Z","iopub.status.idle":"2025-05-26T09:47:35.674583Z","shell.execute_reply":"2025-05-26T09:47:35.674341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sid = SliceInferenceDataset(tm, tta=8)\ndl = torch.utils.data.DataLoader(sid, batch_size=1, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.675534Z","iopub.status.idle":"2025-05-26T09:47:35.675942Z","shell.execute_reply":"2025-05-26T09:47:35.675755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pdfs = []\nwhole_dfs = []\nfor ipts, BS, lsli, tta, iid, fnames_raw in tqdm.tqdm(dl):\n    BS, tta, iid, fnames, lsli = BS.item(), tta.item(), iid[0], [e[0] for e in fnames_raw], lsli.item()\n    predicted_ps = []\n    exp_ps = []\n    for i in range(0, lsli, BS):\n    #   ipt = torch.stack([tensor_tfms(cv2.resize(squarify(s[0], 0), (256, 256))) for s in ress[i: BS+i]]).cuda()\n        with torch.no_grad():\n            res = []\n            exp = []\n            for tt in range(tta):\n                ipt = ipts[tt][0].cuda()\n                for model in models:\n                    with torch.cuda.amp.autocast():\n                        ifr = model(ipt, len(ipt))\n                    res.append(ifr[0].float())\n                    exp.append(ifr[1].float())\n        predict_p = [torch.sigmoid(r.cpu()) for r in res]\n        exp_p = [torch.sigmoid(r.cpu()) for r in exp]\n        predict_p = np.stack(predict_p).mean(0)\n        exp_p = np.stack(exp_p).mean(0)\n        predicted_ps.append(predict_p)\n        exp_ps.append(exp_p)\n    p = np.concatenate(predicted_ps)\n    image_df = pd.DataFrame(p, index=fnames)\n    whole_df = pd.DataFrame(np.concatenate(exp_ps).mean(0).reshape(1, 19), index=[iid])\n    whole_dfs.append(whole_df)\n    pdfs.append(image_df) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.67686Z","iopub.status.idle":"2025-05-26T09:47:35.677227Z","shell.execute_reply":"2025-05-26T09:47:35.677051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# tm = tm.reset_index('fname')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.678479Z","iopub.status.idle":"2025-05-26T09:47:35.678877Z","shell.execute_reply":"2025-05-26T09:47:35.67868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_level = pd.concat(whole_dfs)\nimage_pred = image_level.reset_index().merge(\n    tm[['image_id', 'fname']], left_on='index', right_on='image_id', how='left'\n).set_index('fname').drop(['index', 'image_id'], 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.679816Z","iopub.status.idle":"2025-05-26T09:47:35.680208Z","shell.execute_reply":"2025-05-26T09:47:35.68002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pub_pred = pd.concat(pdfs)\nmerge_pred = pub_pred * image_pred.loc[pub_pred.index]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.681209Z","iopub.status.idle":"2025-05-26T09:47:35.681609Z","shell.execute_reply":"2025-05-26T09:47:35.681431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## If any ensemble","metadata":{}},{"cell_type":"code","source":"merge_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.682445Z","iopub.status.idle":"2025-05-26T09:47:35.682793Z","shell.execute_reply":"2025-05-26T09:47:35.68263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ensem = merge_pred + team_pred.loc[merge_pred.index]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.683712Z","iopub.status.idle":"2025-05-26T09:47:35.684108Z","shell.execute_reply":"2025-05-26T09:47:35.683925Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"merge_pred = ensem","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.684935Z","iopub.status.idle":"2025-05-26T09:47:35.685313Z","shell.execute_reply":"2025-05-26T09:47:35.685126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save prediction","metadata":{}},{"cell_type":"code","source":"df = df.set_index('ID')\nmerge_pred.index.name = 'fname'\nmerge_pred = merge_pred.reset_index()\ntm = tm.set_index('fname')\n\nmerge_pred['ID'] = merge_pred['fname'].str.split('_', expand=True)[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.686225Z","iopub.status.idle":"2025-05-26T09:47:35.686663Z","shell.execute_reply":"2025-05-26T09:47:35.686479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"j_pred = []\nfor iid in merge_pred.ID.unique():\n    enc = ''\n    sub_df = merge_pred[merge_pred.ID == iid]\n    for idx, row in sub_df.iterrows():\n        for i in range(19):\n            enc += f'{i} {row[i]} {tm.loc[row.fname].enc} '\n    j_pred.append({\n        'ID': iid,\n        'ImageWidth': df.loc[iid].ImageWidth,\n        'ImageHeight': df.loc[iid].ImageHeight,\n        'PredictionString': enc[:-1]\n    })","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.687653Z","iopub.status.idle":"2025-05-26T09:47:35.688172Z","shell.execute_reply":"2025-05-26T09:47:35.687945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fast_sub = pd.DataFrame(j_pred)\nfast_sub.to_csv('pub.csv')\nfast_sub = fast_sub.set_index('ID')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.689103Z","iopub.status.idle":"2025-05-26T09:47:35.68954Z","shell.execute_reply":"2025-05-26T09:47:35.689288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fast_sub.head(2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.69043Z","iopub.status.idle":"2025-05-26T09:47:35.690869Z","shell.execute_reply":"2025-05-26T09:47:35.690662Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## save","metadata":{}},{"cell_type":"code","source":"sub2 = pd.concat([sample_submission.drop(fast_sub.index), fast_sub], 0)\nsub2 = sub2.loc[sample_submission.index]\nsub2.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T09:47:35.691742Z","iopub.status.idle":"2025-05-26T09:47:35.692153Z","shell.execute_reply":"2025-05-26T09:47:35.691964Z"}},"outputs":[],"execution_count":null}]}