{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":23823,"databundleVersionId":1920183,"sourceType":"competition"},{"sourceId":1885540,"sourceType":"datasetVersion","datasetId":1123128},{"sourceId":1888569,"sourceType":"datasetVersion","datasetId":1125071},{"sourceId":1888577,"sourceType":"datasetVersion","datasetId":1125092},{"sourceId":1983975,"sourceType":"datasetVersion","datasetId":1182793},{"sourceId":8082539,"sourceType":"datasetVersion","datasetId":4770763},{"sourceId":8391068,"sourceType":"datasetVersion","datasetId":4991231},{"sourceId":8391095,"sourceType":"datasetVersion","datasetId":4991254},{"sourceId":8391217,"sourceType":"datasetVersion","datasetId":4991347},{"sourceId":9646947,"sourceType":"datasetVersion","datasetId":5891629},{"sourceId":9689335,"sourceType":"datasetVersion","datasetId":5923436},{"sourceId":9699776,"sourceType":"datasetVersion","datasetId":5931333},{"sourceId":9752812,"sourceType":"datasetVersion","datasetId":5971300},{"sourceId":10251431,"sourceType":"datasetVersion","datasetId":6340877},{"sourceId":10361255,"sourceType":"datasetVersion","datasetId":6417006},{"sourceId":10361680,"sourceType":"datasetVersion","datasetId":6417322},{"sourceId":10493623,"sourceType":"datasetVersion","datasetId":6497116},{"sourceId":10878971,"sourceType":"datasetVersion","datasetId":6759478},{"sourceId":11235549,"sourceType":"datasetVersion","datasetId":7018915},{"sourceId":11305561,"sourceType":"datasetVersion","datasetId":7070334},{"sourceId":11429643,"sourceType":"datasetVersion","datasetId":7158550},{"sourceId":11546068,"sourceType":"datasetVersion","datasetId":7240696},{"sourceId":11558545,"sourceType":"datasetVersion","datasetId":7247340},{"sourceId":11866238,"sourceType":"datasetVersion","datasetId":7456638},{"sourceId":11894071,"sourceType":"datasetVersion","datasetId":7476224},{"sourceId":11917348,"sourceType":"datasetVersion","datasetId":7491891},{"sourceId":11954584,"sourceType":"datasetVersion","datasetId":7516015},{"sourceId":12103282,"sourceType":"datasetVersion","datasetId":7619628},{"sourceId":12105029,"sourceType":"datasetVersion","datasetId":7620936},{"sourceId":12358423,"sourceType":"datasetVersion","datasetId":7791492},{"sourceId":12394488,"sourceType":"datasetVersion","datasetId":7815768},{"sourceId":13136414,"sourceType":"datasetVersion","datasetId":8322265},{"sourceId":13415964,"sourceType":"datasetVersion","datasetId":8514796}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is simple mmdetection infrence script as a base line.\nTraining part can be foud [here](https://www.kaggle.com/its7171/mmdetection-for-segmentation-training).","metadata":{"papermill":{"duration":0.009854,"end_time":"2021-02-02T02:49:13.549001","exception":false,"start_time":"2021-02-02T02:49:13.539147","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install '/kaggle/input/hpacell/albumentations-1.3.0-py3-none-any.whl'\n!pip install '/kaggle/input/env-package/pycocotools-2.0.7-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n!pip install '/kaggle/input/env-package/timm-0.9.16-py3-none-any.whl'\n!pip install '/kaggle/input/env-package/scikit_image-0.19.3-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:48:30.373223Z","iopub.execute_input":"2025-09-25T06:48:30.37346Z","iopub.status.idle":"2025-09-25T06:58:35.874385Z","shell.execute_reply.started":"2025-09-25T06:48:30.373436Z","shell.execute_reply":"2025-09-25T06:58:35.873602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\n\n# 源文件夹路径\nsource_dir = '/kaggle/input/hpacell/hpacellseg-0.1.8/hpacellseg'\n\n# 目标文件夹路径，Kaggle的输出文件夹通常是'/kaggle/working/'\ntarget_dir1 = '/kaggle/working/hpacellseg'\ntarget_dir2 = '/kaggle/working/pytorch_zoo'\ntarget_dir3 = '/kaggle/working/hpapytorch_zoo'\n# 检查目标文件夹是否已存在，如果不存在，则创建\nif not os.path.exists(target_dir1):\n    os.makedirs(target_dir1)\nif not os.path.exists(target_dir2):\n    os.makedirs(target_dir2)\nif not os.path.exists(target_dir3):\n    os.makedirs(target_dir3)\n# 复制整个文件夹\nshutil.copytree(source_dir, target_dir1, dirs_exist_ok=True)\nshutil.copytree('/kaggle/input/hpacell/master', target_dir2, dirs_exist_ok=True)\nshutil.copytree('/kaggle/input/hpapytorchzoozip/pytorch_zoo-master', target_dir3, dirs_exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:35.875517Z","iopub.execute_input":"2025-09-25T06:58:35.875832Z","iopub.status.idle":"2025-09-25T06:58:36.23141Z","shell.execute_reply.started":"2025-09-25T06:58:35.875797Z","shell.execute_reply":"2025-09-25T06:58:36.230768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install '/kaggle/working/hpapytorch_zoo'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:36.233218Z","iopub.execute_input":"2025-09-25T06:58:36.233644Z","iopub.status.idle":"2025-09-25T06:58:41.255943Z","shell.execute_reply.started":"2025-09-25T06:58:36.233624Z","shell.execute_reply":"2025-09-25T06:58:41.255049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install '/kaggle/working/hpacellseg' --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:41.257243Z","iopub.execute_input":"2025-09-25T06:58:41.257556Z","iopub.status.idle":"2025-09-25T06:58:44.397089Z","shell.execute_reply.started":"2025-09-25T06:58:41.257519Z","shell.execute_reply":"2025-09-25T06:58:44.396378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:44.398949Z","iopub.execute_input":"2025-09-25T06:58:44.399704Z","iopub.status.idle":"2025-09-25T06:58:44.403834Z","shell.execute_reply.started":"2025-09-25T06:58:44.399674Z","shell.execute_reply":"2025-09-25T06:58:44.403175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nfrom tqdm.notebook import tqdm\nimport pickle\nfrom itertools import groupby\nimport matplotlib.pyplot as plt\nimport base64\nimport typing as t\nimport zlib\nimport time\nimport random\nimport sys\nimport time\nfrom hpacellseg.utils import label_cell\nfrom hpacellseg.cellsegmentator import CellSegmentator\nimport warnings\nfrom PIL import Image\nfrom torch import nn\nimport timm\nimport torch.nn.functional as F\nfrom torch import nn, Tensor\nfrom torch.nn import MultiheadAttention\nfrom torch.nn.parameter import Parameter\nfrom pycocotools import _mask as coco_mask\nfrom torchvision import transforms\nfrom torchvision.transforms import ToTensor, Normalize\nimport pandas as pd; pd.options.mode.chained_assignment = None;\nimport numpy as np; print(f\"\\t\\t– NUMPY VERSION: {np.__version__}\");\nimport torch\nimport torchvision.transforms as transforms\n# Built In Imports\nfrom collections import Counter\nfrom datetime import datetime\nimport multiprocessing\nfrom glob import glob\nimport warnings\nimport requests\nimport imageio\nimport IPython\nimport urllib\nimport zipfile\nimport pickle\nimport random\nimport shutil\nimport string\nimport math\nimport time\nimport gzip\nimport sys\nimport ast\nimport csv; csv.field_size_limit(sys.maxsize)\nimport io\nimport os\nimport gc\nimport re\nfrom PIL import Image\nimport matplotlib; print(f\"\\t\\t– MATPLOTLIB VERSION: {matplotlib.__version__}\");\nimport plotly\nimport PIL\nimport copy\nimport collections\n# Submission Imports\nfrom pycocotools import _mask as coco_mask\nimport typing as t\nimport base64\nimport zlib\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n\n# PRESETS\nLBL_NAMES = [\"Nucleoplasm\", \"Nuclear Membrane\", \"Nucleoli\", \"Nucleoli Fibrillar Center\", \"Nuclear Speckles\", \"Nuclear Bodies\", \"Endoplasmic Reticulum\", \"Golgi Apparatus\", \"Intermediate Filaments\", \"Actin Filaments\", \"Microtubules\", \"Mitotic Spindle\", \"Centrosome\", \"Plasma Membrane\", \"Mitochondria\", \"Aggresome\", \"Cytosol\", \"Vesicles\", \"Negative\"]\n# INT_2_STR = {x:LBL_NAMES[x] for x in np.arange(19)}\n# INT_2_STR_LOWER = {k:v.lower().replace(\" \", \"_\") for k,v in INT_2_STR.items()}\n# STR_2_INT_LOWER = {v:k for k,v in INT_2_STR_LOWER.items()}\n# STR_2_INT = {v:k for k,v in INT_2_STR.items()}\n# FIG_FONT = dict(family=\"Helvetica, Arial\", size=14, color=\"#7f7f7f\")\n# LABEL_COLORS = [px.colors.label_rgb(px.colors.convert_to_RGB_255(x)) for x in sns.color_palette(\"Spectral\", len(LBL_NAMES))]\n# LABEL_COL_MAP = {str(i):x for i,x in enumerate(LABEL_COLORS)}\n\nprint(\"\\n\\n... IMPORTS COMPLETE ...\\n\")\n\n##### THIS IS FOR PROTOTYPING AND PUBLIC LB PROBING #####\nONLY_PUBLIC = True\n##### THIS IS FOR PROTOTYPING AND PUBLIC LB PROBING#####\n\nif ONLY_PUBLIC:\n    print(\"\\n... ONLY INFERRING ON PUBLIC TEST DATA (USING PRE-PROCESSED DF) ...\\n\")","metadata":{"papermill":{"duration":0.216599,"end_time":"2021-02-02T02:51:44.826551","exception":false,"start_time":"2021-02-02T02:51:44.609952","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:44.404633Z","iopub.execute_input":"2025-09-25T06:58:44.404845Z","iopub.status.idle":"2025-09-25T06:58:54.582094Z","shell.execute_reply.started":"2025-09-25T06:58:44.404826Z","shell.execute_reply":"2025-09-25T06:58:54.581342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# np.bool = np.bool_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:54.582951Z","iopub.execute_input":"2025-09-25T06:58:54.583425Z","iopub.status.idle":"2025-09-25T06:58:54.587274Z","shell.execute_reply.started":"2025-09-25T06:58:54.583398Z","shell.execute_reply":"2025-09-25T06:58:54.586535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define paths to nucleus and cell models for the cellsegmentator class\n# NUC_MODEL = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_nuclei_v1.pth'\n# CELL_MODEL = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_cell_3ch_v1.pth'\n\n\n# Define the path to the competition data directory\nDATA_DIR = \"/kaggle/input/hpa-single-cell-image-classification\"\n\n# Define the paths to the training and testing tfrecord and \n# image folders respectively for the competition data\nTEST_IMG_DIR = os.path.join(DATA_DIR, \"test\")\n\n# Capture all the relevant full image paths for the competition dataset\nTEST_IMG_PATHS = sorted([os.path.join(TEST_IMG_DIR, f_name) for f_name in os.listdir(TEST_IMG_DIR)])\nprint(f\"... The number of testing images is {len(TEST_IMG_PATHS)}\" \\\n      f\"\\n\\t--> i.e. {len(TEST_IMG_PATHS)//4} 4-channel images ...\")\n\n# Define paths to the relevant csv files\nPUB_SS_CSV = \"/kaggle/input/hpa-sample-submission-with-extra-metadata/updated_sample_submission.csv\"\nSWAP_SS_CSV = os.path.join(DATA_DIR, \"sample_submission.csv\")\n\n# Create the relevant dataframe objects\nss_df = pd.read_csv(SWAP_SS_CSV)\n\n# Test Time Augmentation Information\nDO_TTA = False\nTTA_REPEATS = 8\n\n# helps us control whether this is the full submission or just the initial pass\n# IS_DEMO = len(ss_df)==559\nIS_DEMO = False\nprint('IS_DEMO:',IS_DEMO)\nif IS_DEMO:\n    ss_df_1 = ss_df.drop_duplicates(\"ImageWidth\", keep=\"first\")\n    ss_df_2 = ss_df.drop_duplicates(\"ImageWidth\", keep=\"last\")\n    ss_df = pd.concat([ss_df_1, ss_df_2])\n    del ss_df_1; del ss_df_2; gc.collect();\n    print(\"\\n\\nSAMPLE SUBMISSION DATAFRAME\\n\\n\")\n    print('len:',len(ss_df))\n    display(ss_df)\nelse:\n    print(\"\\n\\nSAMPLE SUBMISSION DATAFRAME\\n\\n\")\n    display(ss_df)\n\n# If demo-submission/display we only do a subset of the data\nif ONLY_PUBLIC:\n    pub_ss_df = pd.read_csv(PUB_SS_CSV)\n\n    if IS_DEMO:\n        pub_ss_df_1 = pub_ss_df.drop_duplicates(\"ImageWidth\", keep=\"first\")\n        pub_ss_df_2 = pub_ss_df.drop_duplicates(\"ImageWidth\", keep=\"last\")\n        pub_ss_df = pd.concat([pub_ss_df_1, pub_ss_df_2])\n\n    pub_ss_df.mask_rles = pub_ss_df.mask_rles.apply(lambda x: ast.literal_eval(x))\n    pub_ss_df.mask_bboxes = pub_ss_df.mask_bboxes.apply(lambda x: ast.literal_eval(x))\n    pub_ss_df.mask_sub_rles = pub_ss_df.mask_sub_rles.apply(lambda x: ast.literal_eval(x))\n\n    print(\"\\n\\nTEST DATAFRAME W/ MASKS\\n\\n\")\n    display(pub_ss_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:54.588092Z","iopub.execute_input":"2025-09-25T06:58:54.588398Z","iopub.status.idle":"2025-09-25T06:58:56.453905Z","shell.execute_reply.started":"2025-09-25T06:58:54.588377Z","shell.execute_reply":"2025-09-25T06:58:56.453281Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# help function","metadata":{}},{"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\n\ndef rle_encoding(img, mask_val=1):\n    \"\"\"\n    Turns our masks into RLE encoding to easily store them\n    and feed them into models later on\n    https://en.wikipedia.org/wiki/Run-length_encoding\n    \n    Args:\n        img (np.array): Segmentation array\n        mask_val (int): Which value to use to create the RLE\n        \n    Returns:\n        RLE string\n    \n    \"\"\"\n    dots = np.where(img.T.flatten() == mask_val)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if (b>prev+1): run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n        \n    return ' '.join([str(x) for x in run_lengths])\n\n\ndef rle_to_mask(rle_string, height, width):\n    \"\"\" Convert RLE sttring into a binary mask \n    \n    Args:\n        rle_string (rle_string): Run length encoding containing \n            segmentation mask information\n        height (int): Height of the original image the map comes from\n        width (int): Width of the original image the map comes from\n    \n    Returns:\n        Numpy array of the binary segmentation mask for a given cell\n    \"\"\"\n    rows,cols = height,width\n    rle_numbers = [int(num_string) for num_string in rle_string.split(' ')]\n    rle_pairs = np.array(rle_numbers).reshape(-1,2)\n    img = np.zeros(rows*cols,dtype=np.uint8)\n    for index,length in rle_pairs:\n        index -= 1\n        img[index:index+length] = 255\n    img = img.reshape(cols,rows)\n    img = img.T\n    return img\n\n\ndef create_pred_col(row):\n    \"\"\" Simple function to return the correct prediction string\n    \n    We will want the original public test dataframe submission when it is \n    available. However, we will use the swapped inn submission dataframe\n    when it is not.\n    \n    Args:\n        row (pd.Series): A row in the dataframe\n    \n    Returns:\n        The prediction string\n    \"\"\"\n    if pd.isnull(row.PredictionString_y):\n        return row.PredictionString_x\n    else:\n        return row.PredictionString_y\n    \n    \ndef load_image(img_id, img_dir, testing=False, only_public=False):\n    \"\"\" Load An Image Using ID and Directory Path - Composes 4 Individual Images \"\"\"\n    if only_public:\n        return_axis = -1\n        clr_list = [\"red\", \"green\", \"blue\"]\n    else:\n        return_axis = 0\n        clr_list = [\"red\", \"green\", \"blue\", \"yellow\"]\n    \n    if not testing:\n        rgby = [\n            np.asarray(Image.open(os.path.join(img_dir, img_id+f\"_{c}.png\")), np.uint8) \\\n            for c in [\"red\", \"green\", \"blue\", \"yellow\"]\n        ]\n        return np.stack(rgby, axis=-1)\n    else:\n        # This is for cellsegmentator\n        return np.stack(\n            [np.asarray(decode_img(tf.io.read_file(os.path.join(img_dir, img_id+f\"_{c}.png\")), testing=True), np.uint8)[..., 0] \\\n             for c in clr_list], axis=return_axis,\n        )\n        \n\n\ndef load_img(base_path, image_id):\n    # Define the file paths for each channel\n    red_path = f\"{base_path}/{image_id}_red.png\"\n    green_path = f\"{base_path}/{image_id}_green.png\"\n    blue_path = f\"{base_path}/{image_id}_blue.png\"\n    yellow_path = f\"{base_path}/{image_id}_yellow.png\"\n    # Load each channel image\n    red = cv2.imread(red_path, cv2.IMREAD_UNCHANGED)\n    green = cv2.imread(green_path, cv2.IMREAD_UNCHANGED)\n    blue = cv2.imread(blue_path, cv2.IMREAD_UNCHANGED)\n    yellow = cv2.imread(yellow_path, cv2.IMREAD_UNCHANGED)\n    # Check if any image didn't load properly\n    if any(map(lambda x: x is None, [red, green, blue, yellow])):\n        raise FileNotFoundError(f\"One or more image files of {image_id} could not be found.\")\n    # Merge the channels into a single four-channel image\n    merged_image = cv2.merge([red, green, blue, yellow])\n    # merged_image = cv2.merge([blue, yellow, red, green])\n    return merged_image # shape: (height, width, 4)\n\ndef plot_rgb(arr, figsize=(12,12)):\n    \"\"\" Plot 3 Channel Microscopy Image \"\"\"\n    plt.figure(figsize=figsize)\n    plt.title(f\"RGB Composite Image\", fontweight=\"bold\")\n    plt.imshow(arr)\n    plt.axis(False)\n    plt.show()\n    \n    \ndef convert_rgby_to_rgb(arr):\n    \"\"\" Convert a 4 channel (RGBY) image to a 3 channel RGB image.\n    \n    Advice From Competition Host/User: lnhtrang\n\n    For annotation (by experts) and for the model, I guess we agree that individual \n    channels with full range px values are better. \n    In annotation, we toggled the channels. \n    For visualization purpose only, you can try blending the channels. \n    For example, \n        - red = red + yellow\n        - green = green + yellow/2\n        - blue=blue.\n        \n    Args:\n        arr (numpy array): The RGBY, 4 channel numpy array for a given image\n    \n    Returns:\n        RGB Image\n    \"\"\"\n    \n    rgb_arr = np.zeros_like(arr[..., :-1])\n    rgb_arr[..., 0] = arr[..., 0]\n    rgb_arr[..., 1] = arr[..., 1]+arr[..., 3]/2\n    rgb_arr[..., 2] = arr[..., 2]\n    \n    return rgb_arr\n    \n    \ndef plot_ex(arr, figsize=(20,6), title=None, plot_merged=True, rgb_only=False):\n    \"\"\" Plot 4 Channels Side by Side \"\"\"\n    if plot_merged and not rgb_only:\n        n_images=5 \n    elif plot_merged and rgb_only:\n        n_images=4\n    elif not plot_merged and rgb_only:\n        n_images=4\n    else:\n        n_images=3\n    plt.figure(figsize=figsize)\n    if type(title) == str:\n        plt.suptitle(title, fontsize=20, fontweight=\"bold\")\n\n    for i, c in enumerate([\"Red Channel – Microtubles\", \"Green Channel – Protein of Interest\", \"Blue - Nucleus\", \"Yellow – Endoplasmic Reticulum\"]):\n        if not rgb_only:\n            ch_arr = np.zeros_like(arr[..., :-1])        \n        else:\n            ch_arr = np.zeros_like(arr)\n        if c in [\"Red Channel – Microtubles\", \"Green Channel – Protein of Interest\", \"Blue - Nucleus\"]:\n            ch_arr[..., i] = arr[..., i]\n        else:\n            if rgb_only:\n                continue\n            ch_arr[..., 0] = arr[..., i]\n            ch_arr[..., 1] = arr[..., i]\n        plt.subplot(1,n_images,i+1)\n        plt.title(f\"{c.title()}\", fontweight=\"bold\")\n        plt.imshow(ch_arr)\n        plt.axis(False)\n        \n    if plot_merged:\n        plt.subplot(1,n_images,n_images)\n        \n        if rgb_only:\n            plt.title(f\"Merged RGB\", fontweight=\"bold\")\n            plt.imshow(arr)\n        else:\n            plt.title(f\"Merged RGBY into RGB\", fontweight=\"bold\")\n            plt.imshow(convert_rgby_to_rgb(arr))\n        plt.axis(False)\n        \n    plt.tight_layout(rect=[0, 0.2, 1, 0.97])\n    plt.show()\n    \n    \ndef flatten_list_of_lists(l_o_l, to_string=False):\n    if not to_string:\n        return [item for sublist in l_o_l for item in sublist]\n    else:\n        return [str(item) for sublist in l_o_l for item in sublist]\n\n\ndef create_segmentation_maps(list_of_image_lists, segmentator, batch_size=8):\n    \"\"\" Function to generate segmentation maps using CellSegmentator tool \n    \n    Args:\n        list_of_image_lists (list of lists):\n            - [[micro-tubules(red)], [endoplasmic-reticulum(yellow)], [nucleus(blue)]]\n        batch_size (int): Batch size to use in generating the segmentation masks\n        \n    Returns:\n        List of lists containing RLEs for all the cells in all images\n    \"\"\"\n    \n    all_mask_rles = {}\n    for i in tqdm(range(0, len(list_of_image_lists[0]), batch_size), total=len(list_of_image_lists[0])//batch_size):\n        \n        # Get batch of images\n        sub_images = [img_channel_list[i:i+batch_size] for img_channel_list in list_of_image_lists] # 0.000001 seconds\n\n        # Do segmentation\n        cell_segmentations = segmentator.pred_cells(sub_images)\n        nuc_segmentations = segmentator.pred_nuclei(sub_images[2])\n\n        # post-processing\n        for j, path in enumerate(sub_images[0]):\n            img_id = path.replace(\"_red.png\", \"\").rsplit(\"/\", 1)[1]\n            nuc_mask, cell_mask = label_cell(nuc_segmentations[j], cell_segmentations[j])\n            new_name = os.path.basename(path).replace('red','mask')\n            all_mask_rles[img_id] = [rle_encoding(cell_mask, mask_val=k) for k in range(1, np.max(cell_mask)+1)]\n    return all_mask_rles\n\n\ndef get_img_list(img_dir, return_ids=False, sub_n=None):\n    \"\"\" Get image list in the format expected by the CellSegmentator tool \"\"\"\n    if sub_n is None:\n        sub_n=len(glob(img_dir + '/' + f'*_red.png'))\n    if return_ids:\n        images = [sorted(glob(img_dir + '/' + f'*_{c}.png'))[:sub_n] for c in [\"red\", \"yellow\", \"blue\"]]\n        return [x.replace(\"_red.png\", \"\").rsplit(\"/\", 1)[1] for x in images[0]], images\n    else:\n        return [sorted(glob(img_dir + '/' + f'*_{c}.png'))[:sub_n] for c in [\"red\", \"yellow\", \"blue\"]]\n    \n    \ndef get_contour_bbox_from_rle(rle, width, height, return_mask=True,):\n    \"\"\" Get bbox of contour as `xmin ymin xmax ymax`\n    \n    Args:\n        rle (rle_string): Run length encoding containing \n            segmentation mask information\n        height (int): Height of the original image the map comes from\n        width (int): Width of the original image the map comes from\n    \n    Returns:\n        Numpy array for a cell bounding box coordinates\n    \"\"\"\n    mask = rle_to_mask(rle, height, width).copy()\n    cnts = grab_contours(\n        cv2.findContours(\n            mask, \n            cv2.RETR_EXTERNAL, \n            cv2.CHAIN_APPROX_SIMPLE\n        ))\n    x,y,w,h = cv2.boundingRect(cnts[0])\n    \n    if return_mask:\n        return (x,y,x+w,y+h), mask\n    else:\n        return (x,y,x+w,y+h)\n    \n\ndef get_contour_bbox_from_raw(raw_mask):\n    \"\"\" Get bbox of contour as `xmin ymin xmax ymax`\n    \n    Args:\n        raw_mask (nparray): Numpy array containing segmentation mask information\n    \n    Returns:\n        Numpy array for a cell bounding box coordinates\n    \"\"\"\n    cnts = grab_contours(\n        cv2.findContours(\n            raw_mask, \n            cv2.RETR_EXTERNAL, \n            cv2.CHAIN_APPROX_SIMPLE\n        ))\n    xywhs = [cv2.boundingRect(cnt) for cnt in cnts]\n    xys = [(xywh[0], xywh[1], xywh[0]+xywh[2], xywh[1]+xywh[3]) for xywh in xywhs]\n    return sorted(xys, key=lambda x: (x[1], x[0]))\n\n\ndef pad_to_square(a):\n    \"\"\" Pad an array `a` evenly until it is a square \"\"\"\n    if a.shape[1]>a.shape[0]: # pad height\n        n_to_add = a.shape[1]-a.shape[0]\n        top_pad = n_to_add//2\n        bottom_pad = n_to_add-top_pad\n        a = np.pad(a, [(top_pad, bottom_pad), (0, 0), (0, 0)], mode='constant')\n\n    elif a.shape[0]>a.shape[1]: # pad width\n        n_to_add = a.shape[0]-a.shape[1]\n        left_pad = n_to_add//2\n        right_pad = n_to_add-left_pad\n        a = np.pad(a, [(0, 0), (left_pad, right_pad), (0, 0)], mode='constant')\n    else:\n        pass\n    return a\n\n\ndef cut_out_cells(rgby, rles, resize_to=(256,256), square_off=True, return_masks=False, from_raw=True):\n    \"\"\" Cut out the cells as padded square images \n    \n    Args:\n        rgby (np.array): 4 Channel image to be cut into tiles\n        rles (list of RLE strings): List of run length encoding containing \n            segmentation mask information\n        resize_to (tuple of ints, optional): The square dimension to resize the image to\n        square_off (bool, optional): Whether to pad the image to a square or not\n        \n    Returns:\n        list of square arrays representing squared off cell images\n    \"\"\"\n    w,h = rgby.shape[:2]\n    contour_bboxes = [get_contour_bbox(rle, w, h, return_mask=return_masks) for rle in rles]\n    if return_masks:\n        masks = [x[-1] for x in contour_bboxes]\n        contour_bboxes = [x[:-1] for x in contour_bboxes]\n    \n    arrs = [rgby[bbox[1]:bbox[3], bbox[0]:bbox[2], ...] for bbox in contour_bboxes]\n    if square_off:\n        arrs = [pad_to_square(arr) for arr in arrs]\n        \n    if resize_to is not None:\n        arrs = [\n            cv2.resize(pad_to_square(arr).astype(np.float32), \n                       resize_to, \n                       interpolation=cv2.INTER_CUBIC) \\\n            for arr in arrs\n        ]\n    if return_masks:\n        return arrs, masks\n    else:\n        return arrs\n\n\ndef grab_contours(cnts):\n    # if the length the contours tuple returned by cv2.findContours\n    # is '2' then we are using either OpenCV v2.4, v4-beta, or\n    # v4-official\n    if len(cnts) == 2:\n        cnts = cnts[0]\n\n    # if the length of the contours tuple is '3' then we are using\n    # either OpenCV v3, v4-pre, or v4-alpha\n    elif len(cnts) == 3:\n        cnts = cnts[1]\n\n    # otherwise OpenCV has changed their cv2.findContours return\n    # signature yet again and I have no idea WTH is going on\n    else:\n        raise Exception((\"Contours tuple must have length 2 or 3, \"\n            \"otherwise OpenCV changed their cv2.findContours return \"\n            \"signature yet again. Refer to OpenCV's documentation \"\n            \"in that case\"))\n\n    # return the actual contours array\n    return cnts\n\n\ndef plot_predictions(img, masks, preds, confs=None, fill_alpha=0.3, lbl_as_str=True):\n    # Initialize\n    FONT = cv2.FONT_HERSHEY_SIMPLEX; FONT_SCALE = 0.7; FONT_THICKNESS = 2; FONT_LINE_TYPE = cv2.LINE_AA;\n    COLORS = [[round(y*255) for y in x] for x in sns.color_palette(\"Spectral\", len(LBL_NAMES))]\n    to_plot = img.copy()\n    cntr_img = img.copy()\n    if confs==None:\n        confs = [None,]*len(masks)\n\n    cnts = grab_contours(\n        cv2.findContours(\n            masks, \n            cv2.RETR_EXTERNAL, \n            cv2.CHAIN_APPROX_SIMPLE\n        ))\n    cnts = sorted(cnts, key=lambda x: (cv2.boundingRect(x)[1], cv2.boundingRect(x)[0]))\n        \n    for c, pred, conf in zip(cnts, preds, confs):\n        # We can only display one color so we pick the first\n        color = COLORS[pred[0]]\n        if not lbl_as_str:\n            classes = \"CLS=[\"+\",\".join([str(p) for p in pred])+\"]\"\n        else:\n            classes = \", \".join([INT_2_STR[p] for p in pred])\n        M = cv2.moments(c)\n        cx = int(M['m10']/M['m00'])\n        cy = int(M['m01']/M['m00'])\n        \n        text_width, text_height = cv2.getTextSize(classes, FONT, FONT_SCALE, FONT_THICKNESS)[0]\n        \n        # Border and fill\n        cv2.drawContours(to_plot, [c], contourIdx=-1, color=[max(0, x-40) for x in color], thickness=10)\n        cv2.drawContours(cntr_img, [c], contourIdx=-1, color=(color), thickness=-1)\n        \n        # Text\n        cv2.putText(to_plot, classes, (cx-text_width//2,cy-text_height//2),\n                    FONT, FONT_SCALE, [min(255, x+40) for x in color], FONT_THICKNESS, FONT_LINE_TYPE)\n    \n    cv2.addWeighted(cntr_img, fill_alpha, to_plot, 1-fill_alpha, 0, to_plot)\n    plt.figure(figsize=(16,16))\n    plt.imshow(to_plot)\n    plt.axis(False)\n    plt.show()\n    \n\n\ndef split_cells(image, mask, img_size):\n    \"\"\"\n    Extracts cells , pad to square and resize to img_size(256*256)\n\n    Args:\n    image (numpy.ndarray): The four-channel image from which to extract cells.\n    mask (numpy.ndarray): The mask image where each cell is labeled with a unique non-zero integer.\n\n    Returns:\n    list of cell image: contains a cell image\n    \"\"\"\n\n    def 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)\n\n    # 确保image和mask的形状匹配\n    assert image.shape[:2] == mask.shape, \"Image and mask shapes do not match\"\n\n    # 获取唯一的细胞标签（排除背景0）\n    unique_cells = np.unique(mask)\n    cell_labels = unique_cells[unique_cells != 0]  # 排除背景 (0)\n    img_w, img_h = image.shape[:2]\n    extracted_cells = []\n    # plt.imshow(mask, cmap='gray')\n    # plt.show()\n    for label in cell_labels:\n        # 创建当前细胞的二值掩码\n        cell_mask = (mask == label).astype(np.uint8)\n        # 找到细胞的边界框\n        rows, cols = np.where(cell_mask)\n        top, bottom, left, right = rows.min(), rows.max(), cols.min(), cols.max()\n        # 提取细胞区域\n        cell_image = image[top:bottom + 1, left:right + 1].copy()\n        cell_mask = cell_mask[top:bottom + 1, left:right + 1]\n        cell_mask = np.repeat(cell_mask[:, :, np.newaxis], 4, axis=2)\n        cell_image = cell_image * cell_mask\n\n        # Pad to square\n        cell_image = squarify(cell_image, 0)\n\n        # Resize to img_size\n        # cell_image = cv2.resize(cell_image, img_size, interpolation=cv2.INTER_LINEAR).astype(np.uint8)\n        cell_image = cv2.resize(cell_image, img_size)\n\n        # 可视化img\n        # plt.imshow(cell_image)\n        # plt.show()\n\n        extracted_cells.append(cell_image)\n\n    return extracted_cells","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.456256Z","iopub.execute_input":"2025-09-25T06:58:56.456488Z","iopub.status.idle":"2025-09-25T06:58:56.500548Z","shell.execute_reply.started":"2025-09-25T06:58:56.456469Z","shell.execute_reply":"2025-09-25T06:58:56.499438Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ","metadata":{"papermill":{"duration":0.016707,"end_time":"2021-02-02T02:51:44.860458","exception":false,"start_time":"2021-02-02T02:51:44.843751","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# model declare and build","metadata":{}},{"cell_type":"code","source":"class ResNestXdual(nn.Module):\n    def __init__(self, pretrained=False):\n        super(ResNestXdual, self).__init__()\n        self.model = timm.create_model(\"resnest50d_4s2x40d\", pretrained=pretrained, in_chans=4)\n        # print(self.model.default_cfg)\n        # self.model.conv1[0].in_channels = 4\n        # self.model.conv1[0].weight = torch.nn.Parameter(\n        #     torch.cat(\n        #         [\n        #             self.model.conv1[0].weight,\n        #             self.model.conv1[0].weight[:, 0:1, :, :],\n        #         ],\n        #         axis=1,\n        #     )\n        # )\n        self.local_fc = nn.Linear(2048, out_features=19)\n        self.global_fc = nn.Linear(2048, out_features=19)\n        self.pool = Actnet(2.5, 2.2)\n        self.dropout = nn.Dropout(0.5)\n\n    def forward(self, x, cnt=16):\n        x = self.model.forward_features(x)\n        desc = nn.Flatten()(self.pool(x))\n        desc_sh = desc.view(-1, cnt, desc.shape[-1])\n        desc_sh = desc_sh.max(1)[0]\n        local_pred = self.local_fc(self.dropout(desc))\n        global_pred = self.global_fc(self.dropout(desc_sh))\n        return local_pred, global_pred\n\n\ndef actnet(x, p, q, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(\n        1.0 / q\n    )\n\n\nclass Actnet(nn.Module):\n    def __init__(self, p, q, eps=1e-6):\n        super(Actnet, self).__init__()\n        self.p = Parameter(torch.ones(1) * p)\n        self.q = Parameter(torch.ones(1) * q)\n        self.eps = eps\n\n    def forward(self, x):\n        return actnet(x, p=self.p, q=self.q, eps=self.eps)\n\n    def __repr__(self):\n        return (\n                self.__class__.__name__\n                + \"(\"\n                + \"p=\"\n                + \"{:.4f}\".format(self.p.data.tolist()[0])\n                + \", \"\n                + \"q=\"\n                + \"{:.4f}\".format(self.q.data.tolist()[0])\n                + \",\"\n                + \"eps=\"\n                + str(self.eps)\n                + \")\"\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.501668Z","iopub.execute_input":"2025-09-25T06:58:56.501978Z","iopub.status.idle":"2025-09-25T06:58:56.521548Z","shell.execute_reply.started":"2025-09-25T06:58:56.501957Z","shell.execute_reply":"2025-09-25T06:58:56.520576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import Optional\nfrom torch import Tensor\nclass TransformerEncoder(nn.Module):\n\n    def __init__(self, d_model, nhead, num_layers, norm=None):\n        super().__init__()\n        encoder_layer = TransformerEncoderLayer(d_model=d_model, nhead=nhead, dim_feedforward=2048,\n                                                dropout=0.1, activation=\"relu\", normalize_before=False)\n        self.layers = _get_clones(encoder_layer, num_layers)\n        self.num_layers = num_layers\n        self.norm = norm\n\n    def forward(self, src,\n                mask: Optional[Tensor] = None,\n                src_key_padding_mask: Optional[Tensor] = None,\n                pos: Optional[Tensor] = None):\n        output = src\n\n        for layer in self.layers:\n            output = layer(output, src_mask=mask,\n                           src_key_padding_mask=src_key_padding_mask, pos=pos)\n\n        if self.norm is not None:\n            output = self.norm(output)\n\n        return output\n\n\nclass TransformerDecoder(nn.Module):\n\n    def __init__(self, d_model=512, nhead=8, num_layers=6, dropout=0.1, return_intermediate=False):\n        super().__init__()\n        decoder_layer = TransformerDecoderLayer(d_model, nhead, dim_feedforward=2048,\n                                                dropout=dropout, activation=\"relu\", normalize_before=False)\n        self.layers = _get_clones(decoder_layer, num_layers)\n        self.num_layers = num_layers\n        self.norm = nn.LayerNorm(d_model)\n        self.return_intermediate = return_intermediate\n\n        self._reset_parameters()\n\n        self.d_model = d_model\n        self.nhead = nhead\n\n    def _reset_parameters(self):\n        for p in self.parameters():\n            if p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n\n    def forward(self, tgt, memory,\n                tgt_mask: Optional[Tensor] = None,\n                memory_mask: Optional[Tensor] = None,\n                tgt_key_padding_mask: Optional[Tensor] = None,\n                memory_key_padding_mask: Optional[Tensor] = None,\n                pos: Optional[Tensor] = None,\n                query_pos: Optional[Tensor] = None):\n        output = tgt\n\n        intermediate = []\n        for layer in self.layers:\n            output  = layer(output, memory, tgt_mask=tgt_mask,\n                           memory_mask=memory_mask,\n                           tgt_key_padding_mask=tgt_key_padding_mask,\n                           memory_key_padding_mask=memory_key_padding_mask,\n                           pos=pos, query_pos=query_pos)\n            if self.return_intermediate:\n                intermediate.append(self.norm(output))\n\n        if self.norm is not None:\n            output = self.norm(output)\n            if self.return_intermediate:\n                intermediate.pop()\n                intermediate.append(output)\n\n        if self.return_intermediate:\n            return torch.stack(intermediate)\n\n        return output.unsqueeze(0)\n\n\nclass TransformerEncoderLayer(nn.Module):\n\n    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1,\n                 activation=\"relu\", normalize_before=False):\n        super().__init__()\n        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout)\n        # Implementation of Feedforward model\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n\n        self.activation = _get_activation_fn(activation)\n        self.normalize_before = normalize_before\n\n        self.debug_mode = False\n        self.debug_name = None\n\n    def with_pos_embed(self, tensor, pos: Optional[Tensor]):\n        return tensor if pos is None else tensor + pos\n\n    def forward_post(self,\n                     src,\n                     src_mask: Optional[Tensor] = None,\n                     src_key_padding_mask: Optional[Tensor] = None,\n                     pos: Optional[Tensor] = None):\n        q = k = self.with_pos_embed(src, pos)\n        src2, corr = self.self_attn(q, k, value=src, attn_mask=src_mask,\n                                    key_padding_mask=src_key_padding_mask)\n\n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))\n        src = src + self.dropout2(src2)\n        src = self.norm2(src)\n        return src\n\n    def forward_pre(self, src,\n                    src_mask: Optional[Tensor] = None,\n                    src_key_padding_mask: Optional[Tensor] = None,\n                    pos: Optional[Tensor] = None):\n        src2 = self.norm1(src)\n        q = k = self.with_pos_embed(src2, pos)\n        src2 = self.self_attn(q, k, value=src2, attn_mask=src_mask,\n                              key_padding_mask=src_key_padding_mask)[0]\n\n        src = src + self.dropout1(src2)\n        src2 = self.norm2(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src2))))\n        src = src + self.dropout2(src2)\n        return src\n\n    def forward(self, src,\n                src_mask: Optional[Tensor] = None,\n                src_key_padding_mask: Optional[Tensor] = None,\n                pos: Optional[Tensor] = None):\n        if self.normalize_before:\n            return self.forward_pre(src, src_mask, src_key_padding_mask, pos)\n        return self.forward_post(src, src_mask, src_key_padding_mask, pos)\n\n\nclass TransformerDecoderLayer(nn.Module):\n\n    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1,\n                 activation=\"relu\", normalize_before=False):\n        super().__init__()\n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        # Implementation of Feedforward model\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.norm3 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.dropout3 = nn.Dropout(dropout)\n\n        self.activation = _get_activation_fn(activation)\n        self.normalize_before = normalize_before\n\n    def with_pos_embed(self, tensor, pos: Optional[Tensor]):\n        return tensor if pos is None else tensor + pos\n\n    def forward_post(self, tgt, memory,\n                     tgt_mask: Optional[Tensor] = None,\n                     memory_mask: Optional[Tensor] = None,\n                     tgt_key_padding_mask: Optional[Tensor] = None,\n                     memory_key_padding_mask: Optional[Tensor] = None,\n                     pos: Optional[Tensor] = None,\n                     query_pos: Optional[Tensor] = None):\n        q = k = self.with_pos_embed(tgt, query_pos)\n        tgt2, self_attns = self.self_attn(q, k, value=tgt, attn_mask=tgt_mask,\n                              key_padding_mask=tgt_key_padding_mask)\n        tgt = tgt + self.dropout1(tgt2)\n        tgt = self.norm1(tgt)\n        tgt2, cross_attns = self.multihead_attn(query=self.with_pos_embed(tgt, query_pos),\n                                   key=self.with_pos_embed(memory, pos),\n                                   value=memory, attn_mask=memory_mask,\n                                   key_padding_mask=memory_key_padding_mask)\n        tgt = tgt + self.dropout2(tgt2)\n        tgt = self.norm2(tgt)\n        tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))\n        tgt = tgt + self.dropout3(tgt2)\n        tgt = self.norm3(tgt)\n        return tgt\n\n    def forward_pre(self, tgt, memory,\n                    tgt_mask: Optional[Tensor] = None,\n                    memory_mask: Optional[Tensor] = None,\n                    tgt_key_padding_mask: Optional[Tensor] = None,\n                    memory_key_padding_mask: Optional[Tensor] = None,\n                    pos: Optional[Tensor] = None,\n                    query_pos: Optional[Tensor] = None):\n        tgt2 = self.norm1(tgt)\n        q = k = self.with_pos_embed(tgt2, query_pos)\n        tgt2 = self.self_attn(q, k, value=tgt2, attn_mask=tgt_mask,\n                              key_padding_mask=tgt_key_padding_mask)[0]\n        tgt = tgt + self.dropout1(tgt2)\n        tgt2 = self.norm2(tgt)\n        tgt2 = self.multihead_attn(query=self.with_pos_embed(tgt2, query_pos),\n                                   key=self.with_pos_embed(memory, pos),\n                                   value=memory, attn_mask=memory_mask,\n                                   key_padding_mask=memory_key_padding_mask)[0]\n        tgt = tgt + self.dropout2(tgt2)\n        tgt2 = self.norm3(tgt)\n        tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))\n        tgt = tgt + self.dropout3(tgt2)\n        return tgt\n\n    def forward(self, tgt, memory,\n                tgt_mask: Optional[Tensor] = None,\n                memory_mask: Optional[Tensor] = None,\n                tgt_key_padding_mask: Optional[Tensor] = None,\n                memory_key_padding_mask: Optional[Tensor] = None,\n                pos: Optional[Tensor] = None,\n                query_pos: Optional[Tensor] = None):\n        if self.normalize_before:\n            return self.forward_pre(tgt, memory, tgt_mask, memory_mask,\n                                    tgt_key_padding_mask, memory_key_padding_mask, pos, query_pos)\n        return self.forward_post(tgt, memory, tgt_mask, memory_mask,\n                                 tgt_key_padding_mask, memory_key_padding_mask, pos, query_pos)\n\n\ndef _get_clones(module, N):\n    return nn.ModuleList([copy.deepcopy(module) for i in range(N)])\n\n\ndef _get_activation_fn(activation):\n    \"\"\"Return an activation function given a string\"\"\"\n    if activation == \"relu\":\n        return F.relu\n    if activation == \"gelu\":\n        return F.gelu\n    if activation == \"glu\":\n        return F.glu\n    raise RuntimeError(F\"activation should be relu/gelu, not {activation}.\")\n\n\n\nclass Backbone(nn.Module):\n    def __init__(self, model_name='resnet50d', in_chans=4, pretrained=False, cfg_overlay=None):\n        super().__init__()\n        self.model_name = model_name\n        if model_name in ['resnet50d', 'resnet101d', 'resnet200d']:\n            self.model = timm.create_model(model_name, pretrained=pretrained, in_chans=in_chans,\n                                           pretrained_cfg_overlay=cfg_overlay, num_classes=0)\n            self.dim_feats = 2048\n        if model_name == 'vit':\n            self.model = timm.create_model('vit_medium_patch16_reg4_gap_256.sbb_in12k_ft_in1k',\n                                           features_only=True,\n                                           pretrained=pretrained, in_chans=in_chans)\n            self.dim_feats = 512\n        if model_name == 'convnextv2':\n            self.model = timm.create_model('convnextv2_base.fcmae_ft_in22k_in1k_384', pretrained=pretrained,\n                                           in_chans=in_chans,\n                                           pretrained_cfg_overlay=cfg_overlay)\n            self.dim_feats = 1024\n            # (bs, 1024, 12, 12)\n\n    def forward(self, x):\n        if self.model_name == 'vit':\n            return self.model(x)[-1]\n\n        return self.model.forward_features(x)\n\n\n# Disambiguation Multi-label Learning (DML) with Multiple Instance Learning (MIL)\nclass DML_MIL2(nn.Module):\n    def __init__(self, backbone='resnet50d', mode='train', out_features=19, pretrained=False,\n                 d_model=512, layers=8, dropout=(0.5, 0.5),\n                 cfg_overlay=None):\n        super().__init__()\n        self.mode = mode\n        self.d_model = d_model\n        # 创建共享的ResNet骨干网络\n        self.backbone = Backbone(backbone, in_chans=4, pretrained=pretrained, cfg_overlay=cfg_overlay)\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n\n        self.cell_dropout = nn.Dropout(dropout[0])\n        self.query_dropout = nn.Dropout(dropout[1])\n\n        self.transformer_decoder = TransformerDecoder(d_model=d_model, nhead=8, num_layers=layers)\n\n        self.fc_cell = nn.Linear(self.backbone.dim_feats, out_features)\n        self.fc_query = GroupWiseLinear(out_features, d_model)\n        # 创建查询向量\n        self.query_tokens = nn.Embedding(out_features, self.d_model)\n\n        self.input_proj = nn.Conv2d(self.backbone.dim_feats, self.d_model, kernel_size=1)\n        self.pos_embedding = PositionEmbeddingLearned(self.d_model // 2)\n\n    # bs_final= bs*cnt\n    def forward(self, cell_images, cnt):\n        # 提取细胞图像的特征\n        cell_features = self.backbone(cell_images)  # [bs_final, 2048, 7, 7]\n        if cell_features.shape[1] != self.d_model:\n            memory = self.input_proj(cell_features)  # [bs_final, d_model, 7, 7]\n        else:\n            memory = cell_features\n        bs_final, emb_dim, h, w = memory.shape\n        bs = bs_final // cnt\n        pos_emb = self.pos_embedding(memory)  # [bs_final, d_model, 7, 7]\n\n        # memory = memory.flatten(2).permute(0, 2, 1).contiguous().view(-1, h * w * cnt, emb_dim)\n        memory = memory.flatten(2).permute(0, 2, 1).contiguous().view(bs, cnt, h * w,\n                                                                      emb_dim)  # [bs, cnt, h*w, d_model]\n        # cell_pos_emb = self.cell_pos_embedding(memory, cnt)  # [1, cnt, 1, d_model]\n        # memory = memory + cell_pos_emb  # [bs, cnt, h*w, d_model]\n        memory = memory.view(bs, h * w * cnt, emb_dim)  # [bs, h*w*cnt, d_model]\n\n        # [bs_final, d_model, 7, 7] ->[bs_final, d_model, 49]-> [bs_final, h*w, d_model] -> [bs, h*w*cnt, d_model]\n        memory = memory.permute(1, 0, 2)  # [49*cnt, bs, d_model]\n        pos_emb = pos_emb.flatten(2).permute(0, 2, 1).contiguous().view(bs, h * w * cnt,\n                                                                        emb_dim)  # [bs, h*w*cnt, d_model]\n        pos_emb = pos_emb.permute(1, 0, 2)  # [49*cnt, bs, d_model]\n\n        # 细胞特征作为key和value\n        # 查询向量作为query\n        query = self.query_tokens.weight.unsqueeze(1).expand(19, bs, -1)  # [19, bs, d_model]\n        # Transformer Decoder\n        trans_output = self.transformer_decoder(query, memory, pos=pos_emb)[0]  # [19, bs, d_model]\n\n        cell_features = nn.Flatten()(self.pooling(cell_features))  # [bs_final, 2048]\n        # 分类预测\n        cell_logits = self.fc_cell(self.cell_dropout(cell_features))  # [bs_final, 19]\n\n        # trans_output = trans_output.permute(1, 0, 2).contiguous().view(-1, emb_dim)  # [bs*19, d_model]\n        # query_logits = self.fc_query(self.query_dropout(trans_output)).view(bs, 19)  # [bs*19, 1] -> [bs, 19]\n        ### group-wise linear\n        trans_output = trans_output[-1].permute(1, 0, 2)  # [bs, 19, d_model]\n        if self.mode == 'debug':\n            return cell_features, trans_output\n        query_logits = self.fc_query(self.query_dropout(trans_output))  # [bs, 19]\n\n        return cell_logits, query_logits\n\n\nclass PositionEmbeddingLearned(nn.Module):\n    \"\"\"\n    Absolute pos embedding, learned.\n    num_pos_feats: feat_dim // 2\n    \"\"\"\n\n    def __init__(self, num_pos_feats=256):\n        super().__init__()\n        self.row_embed = nn.Embedding(50, num_pos_feats)\n        self.col_embed = nn.Embedding(50, num_pos_feats)\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        nn.init.uniform_(self.row_embed.weight)\n        nn.init.uniform_(self.col_embed.weight)\n\n    def forward(self, x):\n        h, w = x.shape[-2:]\n        i = torch.arange(w, device=x.device)\n        j = torch.arange(h, device=x.device)\n        x_emb = self.col_embed(i)\n        y_emb = self.row_embed(j)\n        pos = torch.cat([\n            x_emb.unsqueeze(0).repeat(h, 1, 1),\n            y_emb.unsqueeze(1).repeat(1, w, 1),\n        ], dim=-1).permute(2, 0, 1).unsqueeze(0).repeat(x.shape[0], 1, 1, 1)\n        return pos\n\n\nclass CellPositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_cells):\n        \"\"\"\n        d_model: 特征维度\n        max_cells: 最大细胞数量 cnt\n        \"\"\"\n        super(CellPositionalEncoding, self).__init__()\n\n        # 创建位置编码矩阵\n        pe = torch.zeros(max_cells, d_model)\n        position = torch.arange(0, max_cells, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n\n        pe[:, 0::2] = torch.sin(position * div_term)  # 偶数维度 sin\n        pe[:, 1::2] = torch.cos(position * div_term)  # 奇数维度 cos\n\n        # 添加 batch 和 token 维度\n        pe = pe.unsqueeze(0).unsqueeze(2)  # shape: (1, max_cells, 1, d_model)\n\n        self.register_buffer('pe', pe)\n\n    def forward(self, x, cnt):\n        \"\"\"\n        输入 x 的 shape: (bs, cnt, 49, d_model)\n        cnt: 实际的细胞数量\n        \"\"\"\n        # 提取前 cnt 个细胞的位置编码\n        return self.pe[:, :cnt]  # shape: (1, cnt, 1, d_model)\n\n\nclass GroupWiseLinear(nn.Module):\n    # could be changed to:\n    # output = torch.einsum('ijk,zjk->ij', x, self.W)\n    # or output = torch.einsum('ijk,jk->ij', x, self.W[0])\n    def __init__(self, num_class, hidden_dim, bias=True):\n        super().__init__()\n        self.num_class = num_class\n        self.hidden_dim = hidden_dim\n        self.bias = bias\n\n        self.W = nn.Parameter(torch.Tensor(1, num_class, hidden_dim))\n        if bias:\n            self.b = nn.Parameter(torch.Tensor(1, num_class))\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        stdv = 1. / math.sqrt(self.W.size(2))\n        for i in range(self.num_class):\n            self.W[0][i].data.uniform_(-stdv, stdv)\n        if self.bias:\n            for i in range(self.num_class):\n                self.b[0][i].data.uniform_(-stdv, stdv)\n\n    def forward(self, x):\n        # x: B,K,d\n        x = (self.W * x).sum(-1)\n        if self.bias:\n            x = x + self.b\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.52269Z","iopub.execute_input":"2025-09-25T06:58:56.523001Z","iopub.status.idle":"2025-09-25T06:58:56.567825Z","shell.execute_reply.started":"2025-09-25T06:58:56.522969Z","shell.execute_reply":"2025-09-25T06:58:56.567007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Chanel_MIL(nn.Module):\n    def __init__(self, backbone='resnet50d', mode='train', num_class=19, pretrained=False, d_model=512,\n                 layers=(4, 4),  # encoder_layers, decoder_layers\n                 dropout=(0.5, 0.5),\n                 cfg_overlay=None):\n        super().__init__()\n        self.mode = mode\n        self.d_model = d_model\n        # 创建共享的ResNet骨干网络\n        self.backbone = Backbone(backbone, in_chans=1, pretrained=pretrained, cfg_overlay=cfg_overlay)\n        # self.pooling = nn.AdaptiveAvgPool2d(1)\n\n        self.cell_dropout = nn.Dropout(dropout[0])\n        self.query_dropout = nn.Dropout(dropout[1])\n\n        self.transformer_encoder = TransformerEncoder(d_model=d_model, nhead=8, num_layers=layers[0])\n        self.transformer_decoder = TransformerDecoder(d_model=d_model, nhead=8, num_layers=layers[1])\n\n        self.pooling = nn.AdaptiveAvgPool1d(1)\n        # 定义分类层\n        self.fc_cell = nn.Linear(d_model, num_class)\n        self.fc_query = GroupWiseLinear(num_class, d_model)\n        # 创建查询向量\n        self.query_tokens = nn.Embedding(num_class, self.d_model)\n        # self.query_pos_emb = nn.Embedding(19, self.d_model)\n\n        # self.register_buffer(\"prototypes\", torch.zeros(19, self.d_model))\n\n        self.input_proj = nn.Conv2d(self.backbone.dim_feats, self.d_model, kernel_size=1)\n        num_pos = 4 * 8 * 8\n        if backbone == 'vit':\n            num_pos = 4 * 16 * 16\n        self.pos_embedding = nn.Embedding(num_pos, self.d_model)\n        # self.cell_pos_embedding = CellPositionalEncoding(self.d_model, 100)\n\n        self._reset_parameters()\n\n    def _reset_parameters(self):\n        nn.init.xavier_uniform_(self.fc_cell.weight)\n        nn.init.xavier_uniform_(self.query_tokens.weight)\n        nn.init.xavier_uniform_(self.pos_embedding.weight)\n\n    # x: [bs*cnt,4,h,w]\n    def forward(self, x, cnt):\n        bs_final, _, h, w = x.shape\n        bs = bs_final // cnt\n        x = x.view(-1, h, w)  # [bs*cnt*4, h, w]\n        x = x.unsqueeze(1)  # [bs*cnt*4, 1, h, w]\n        feat = self.backbone(x)  # [bs * cnt * 4, 2048,7,7]\n        h0, w0 = feat.shape[-2], feat.shape[-1]\n        if feat.shape[1] != self.d_model:\n            src = self.input_proj(feat)  # [bs * cnt * 4, d_model,7,7]\n        else:\n            src = feat\n        src = src.view(bs * cnt, 4, self.d_model, h0, w0)  # [bs * cnt , 4, d_model,7,7]\n\n        src = src.flatten(3).permute(0, 1, 3, 2).contiguous().view(bs * cnt, 4 * h0 * w0, self.d_model)\n        # -> [bs * cnt , 4, d_model,49]->[bs * cnt , 4, 49, d_model] -> [bs * cnt , 4*49, d_model]\n        pos = self.pos_embedding.weight.repeat(bs_final, 1, 1)  # [bs_final, 4*49, d_model]\n        memery = self.transformer_encoder(src.permute(1, 0, 2), pos=pos.permute(1, 0, 2))  # [4*49, bs*cnt, d_model]\n        memery = memery.permute(1, 0, 2).contiguous()  # [bs*cnt, 4*49, d_model]\n        feature = self.pooling(memery.transpose(1, 2)).flatten(1)  # [bs*cnt, d_model]\n        cell_logits = self.fc_cell(self.cell_dropout(feature))  # [bs*cnt, 19]\n\n        # memery = memery + pos\n        memery = memery.view(bs, cnt, 4 * h0 * w0, self.d_model)  # [bs, cnt, 4*49, d_model]\n        # cell_pos = self.cell_pos_embedding(cnt)  # [1, cnt, 1, d_model]\n        # memery = memery + cell_pos  # [bs, cnt, 4*49, d_model]\n        memery = memery.view(bs, cnt * 4 * h0 * w0, self.d_model)  # [bs, cnt*4*49, d_model]\n\n        query = self.query_tokens.weight.unsqueeze(1).expand(19, bs, -1)  # [19, bs, d_model]\n        out, self_attns, cross_attns = self.transformer_decoder(query, memery.permute(1, 0, 2))  # [19, bs, d_model]\n        query_logits = self.fc_query(self.query_dropout(out[-1].permute(1, 0, 2)))  # [bs, 19]\n\n        return cell_logits, query_logits\n\nclass Cellnet(nn.Module):\n    def __init__(self, model_name='resnet50d', in_chan_num=4, dropout=0.5, class_num=19, pretrained=False):\n        super(Cellnet, self).__init__()\n        self.backbone = Backbone(model_name=model_name, in_chans=in_chan_num, pretrained=pretrained)\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.dropout = nn.Dropout(dropout)\n        self.linear = nn.Linear(self.backbone.dim_feats, class_num)\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.pooling(x)\n        x = x.view(x.size(0), -1)\n        logits = self.linear(self.dropout(x))\n        return logits\n\nclass LC_MIL(nn.Module):\n    def __init__(self, backbone='resnet50d', mode='train', out_features=19, pretrained=False,\n                 d_model=512, layers=8, dropout=(0.5, 0.5),\n                 cfg_overlay=None):\n        super().__init__()\n        self.mode = mode\n        self.d_model = d_model\n        # 创建共享的ResNet骨干网络\n        self.backbone = Backbone(backbone, in_chans=4, pretrained=pretrained, cfg_overlay=cfg_overlay)\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n\n        self.cell_dropout = nn.Dropout(dropout[0])\n        self.image_dropout = nn.Dropout(dropout[1])\n\n        self.transformer_decoder = TransformerDecoder(d_model=d_model, nhead=8, num_layers=layers)\n\n        # 定义分类层\n        self.fc_cell = nn.Linear(self.backbone.dim_feats, out_features)\n        self.fc_query = GroupWiseLinear(out_features, d_model)\n        # 创建查询向量\n        self.query_tokens = nn.Embedding(out_features, self.d_model)\n\n        self.input_proj = nn.Conv2d(self.backbone.dim_feats, self.d_model, kernel_size=1)\n        self.pos_embedding = PositionEmbeddingLearned(self.d_model // 2)\n\n    # bs_final= bs*cnt\n    def forward(self, cell_images, cnt=16):\n        # 提取细胞图像的特征\n        cell_features = self.backbone(cell_images)  # [bs_final, 2048, 7, 7]\n        if cell_features.shape[1] != self.d_model:\n            memory = self.input_proj(cell_features)  # [bs_final, d_model, 7, 7]\n        else:\n            memory = cell_features\n        bs_final, emb_dim, h, w = memory.shape\n        bs = bs_final // cnt\n        pos_emb = self.pos_embedding(memory)  # [bs_final, d_model, 7, 7]\n        # -> [bs_final, d_model, 49] -> [49, bs_final, d_model]\n        memory = memory.flatten(2).permute(2, 0, 1)\n        pos_emb = pos_emb.flatten(2).permute(2, 0, 1)  # [49, bs_final, d_model]\n\n        # 细胞特征作为key和value\n        # 查询向量作为query\n        query = self.query_tokens.weight.unsqueeze(1).expand(19, bs_final, emb_dim)  # [19, bs_final, d_model]\n        # Transformer Decoder\n        trans_output, self_attns, cross_attns = self.transformer_decoder(query, memory,\n                                                                         pos=pos_emb)  # list[19, bs_final, d_model]\n        # cross_attns: list[bs_final,19, h*w]\n        # self_attns: list[bs_final,19, 19]\n        # group-wise linear\n        local_features = trans_output[-1].permute(1, 0, 2)  # [bs_final, 19, d_model]\n        global_features = local_features.view(bs, cnt, 19, self.d_model)  # [bs, cnt, 19, d_model]\n        global_features = global_features.max(1)[0]  # [bs, 19, d_model]\n        global_logits = self.fc_query(self.image_dropout(global_features))  # [bs, 19]\n        local_logits = self.fc_query(self.cell_dropout(local_features))  # [bs_final, 19]\n        if self.mode == 'feature':\n            return cell_features\n        if self.mode == 'debug':\n            return local_logits, global_logits, self_attns, cross_attns\n        if self.mode == 'test':\n            return local_logits, global_logits\n\n        return local_logits, global_logits, local_features, global_features\nclass LC_MIL2(nn.Module):\n    def __init__(\n            self,\n            backbone=\"resnet50d\",\n            mode=\"train\",\n            out_features=19,\n            pretrained=False,\n            d_model=512,\n            prev_dim=256,\n            layers=8,\n            dropout=0.5,\n            cfg_overlay=None,\n    ):\n        super().__init__()\n        self.mode = mode\n        self.out_features = out_features\n        self.d_model = d_model\n        # 创建共享的ResNet骨干网络\n        self.backbone = Backbone(\n            backbone, in_chans=4, pretrained=pretrained, cfg_overlay=cfg_overlay\n        )\n        # self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.cell_dropout = nn.Dropout(dropout)\n        # self.image_dropout = nn.Dropout(dropout[1])\n\n        self.transformer_decoder = TransformerDecoder(\n            d_model=d_model, nhead=8, num_layers=layers\n        )\n\n        # 定义分类层\n        # self.fc_cell = nn.Linear(self.backbone.dim_feats, out_features)\n        self.fc_query = GroupWiseLinear(out_features, d_model)\n        # 创建查询向量\n        self.query_tokens = nn.Embedding(out_features, self.d_model)\n\n        self.input_proj = nn.Conv2d(\n            self.backbone.dim_feats, self.d_model, kernel_size=1\n        )\n        self.pos_embedding = PositionEmbeddingLearned(self.d_model // 2)\n\n        self.projector = nn.Sequential(\n            nn.Linear(d_model, prev_dim, bias=False),\n            nn.BatchNorm1d(prev_dim),\n            nn.ReLU(inplace=True),\n            nn.Linear(prev_dim, prev_dim),\n            nn.BatchNorm1d(prev_dim),\n            nn.ReLU(inplace=True),  # second layer\n            nn.Linear(prev_dim, d_model),\n            nn.BatchNorm1d(d_model, affine=False),  # output layer\n        )\n        # 预测\n        self.predictor = nn.Sequential(\n            nn.Linear(d_model, prev_dim, bias=False),\n            nn.BatchNorm1d(prev_dim),\n            nn.ReLU(inplace=True),  # hidden layer\n            nn.Linear(prev_dim, d_model),\n        )  # output layer\n\n    def _forward_single_branch(self, images):\n        \"\"\"处理单个输入视图 (x1 或 x2) 的前向传播\"\"\"\n        raw_backbone_features = self.backbone(images)  # [bs_final, backbone_dim, H, W]\n        # 投影\n        memory_spatial = self.input_proj(raw_backbone_features)  # [bs_final, d_model, H, W]\n        bs_final, _, h, w = memory_spatial.shape\n        pos_emb_spatial = self.pos_embedding(memory_spatial)  # [bs_final, d_model, H, W]\n        memory_transformer = memory_spatial.flatten(2).permute(2, 0, 1)  # [H*W, bs_final, d_model]\n        pos_emb_transformer = pos_emb_spatial.flatten(2).permute(2, 0, 1)  # [H*W, bs_final, d_model]\n\n        # query_tokens: [out_features, bs_final, d_model]\n        query = self.query_tokens.weight.unsqueeze(1).expand(self.out_features, bs_final, self.d_model)\n\n        # Transformer Decoder  # [out_features, bs_final, d_model]\n        transformer_output = self.transformer_decoder(query, memory_transformer, pos=pos_emb_transformer)\n        # class_tokens 形状: [bs_final, out_features, d_model] - 这些是用于对比学习的特征\n        class_tokens = transformer_output[-1].permute(1, 0, 2)\n\n        return class_tokens\n\n    def forward(self, x1, x2=None):\n        \"\"\"\n        Args:\n            x1: 第一个视图的图像批次 [bs_final, C, H_img, W_img]\n            x2: 第二个视图的图像批次 [bs_final, C, H_img, W_img]\n        \"\"\"\n        # 处理第一个视图\n        class_tokens1 = self._forward_single_branch(x1)\n        # class_tokens1 形状: [bs_final, out_features, d_model]\n\n        # # 7. 计算分类 logits\n        dropped_class_tokens = self.cell_dropout(class_tokens1)\n        local_logits = self.fc_query(dropped_class_tokens)  # [bs_final, out_features]\n\n        if self.mode == \"test\":\n            return local_logits\n\n        # 处理第二个视图 (共享权重)\n        class_tokens2 = self._forward_single_branch(x2)\n        # class_tokens2 形状: [bs_final, out_features, d_model]\n\n        # --- SimSiam 对比学习部分 (作用于Transformer Decoder输出的 class_tokens) ---\n        _bs_final, _num_tokens, _token_dim = class_tokens1.shape  # 获取维度信息\n\n        # 将tokens展平以便输入到Projector/Predictor\n        # 从 [bs_final, num_tokens, token_dim] -> [bs_final * num_tokens, token_dim]\n        tokens1_flat = class_tokens1.reshape(-1, _token_dim)\n        tokens2_flat = class_tokens2.reshape(-1, _token_dim)\n\n        # 通过Projector和Predictor\n        z1_flat = self.projector(tokens1_flat)\n        # [bs_final * num_tokens, simsiam_proj_output_dim]\n        z2_flat = self.projector(tokens2_flat)\n\n        p1_flat = self.predictor(z1_flat)\n        # [bs_final * num_tokens, simsiam_proj_output_dim]\n        p2_flat = self.predictor(z2_flat)\n\n        return p1_flat, p2_flat, z1_flat.detach(), z2_flat.detach(), local_logits\n\nclass Encoder(nn.Module):\n    \"\"\"\n    单独的编码器模块，包含 backbone, 投影, 位置编码, Transformer 解码，和投影头。\n    \"\"\"\n\n    def __init__(\n        self,\n        backbone,\n        d_model,\n        out_features,\n        dropout,\n        pos_dim,\n        transformer_layers,\n        transformer_heads,\n        proj_dim,\n        hid_dim,\n    ):\n        super().__init__()\n        # 特征提取骨干网络\n        self.backbone = backbone\n        # 将 backbone 特征投影到 d_model\n        self.input_proj = nn.Conv2d(backbone.dim_feats, d_model, kernel_size=1)\n        # 位置编码\n        self.pos_embedding = PositionEmbeddingLearned(pos_dim)\n        # Transformer Decoder\n        self.transformer_decoder = TransformerDecoder(\n            d_model=d_model,\n            nhead=transformer_heads,\n            num_layers=transformer_layers,\n        )\n        # 查询向量\n        self.query_tokens = nn.Embedding(out_features, d_model)\n        # 分类 head，不经过 projector\n        self.cell_dropout = nn.Dropout(dropout)\n        self.fc_query = GroupWiseLinear(out_features, d_model)\n        self.projector = nn.Sequential(\n            nn.Linear(d_model, hid_dim),\n            # nn.LayerNorm(proj_dim),\n            nn.ReLU(inplace=True),\n            # nn.Linear(proj_dim, proj_dim),\n            # nn.LayerNorm(proj_dim),\n            # nn.ReLU(inplace=True),\n            nn.Linear(hid_dim, proj_dim),\n            # nn.LayerNorm(d_model, elementwise_affine=False),\n        )\n        # # 在Encoder初始化中添加\n        # for layer in self.projector:\n        #     if isinstance(layer, nn.Linear):\n        #         nn.init.xavier_uniform_(layer.weight)\n        #         nn.init.constant_(layer.bias, 0)\n\n    def forward(self, x):\n        # backbone 特征: [B, C, H, W]\n        feats = self.backbone(x)\n        # 投影 + 位置编码\n        mem = self.input_proj(feats)  # [B, d_model, H, W]\n        pos = self.pos_embedding(mem)  # [B, d_model, H, W]\n        B, _, H, W = mem.shape\n        mem_flat = mem.flatten(2).permute(2, 0, 1)  # [HW, B, d_model]\n        pos_flat = pos.flatten(2).permute(2, 0, 1)  # [HW, B, d_model]\n        # 构建 query: [out_features, B, d_model]\n        q = self.query_tokens.weight.unsqueeze(1).expand(\n            self.query_tokens.num_embeddings, B, self.query_tokens.embedding_dim\n        )\n        # Transformer 解码\n        tgt = self.transformer_decoder(q, mem_flat, pos=pos_flat)\n        tokens = tgt[-1].permute(1, 0, 2).contiguous()  # [B, out_features, d_model]\n        # 返回最终 class tokens: [B, out_features, d_model]\n        logits = self.fc_query(self.cell_dropout(tokens))  # [B, out_features]\n        # tokens = self.projector(tokens)  # [B, out_features, d_model]\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.568855Z","iopub.execute_input":"2025-09-25T06:58:56.569096Z","iopub.status.idle":"2025-09-25T06:58:56.599788Z","shell.execute_reply.started":"2025-09-25T06:58:56.569078Z","shell.execute_reply":"2025-09-25T06:58:56.598982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Parameters\nIMAGE_SIZES = [1728, 2048, 3072, 4096]\nBATCH_SIZE = 1\nCONF_THRESH = 0.0\n\n\n\n# Switch what we will be actually infering on\nif ONLY_PUBLIC:\n    # Make subset dataframes\n    predict_df_1728 = pub_ss_df[pub_ss_df.ImageWidth==IMAGE_SIZES[0]]\n    predict_df_2048 = pub_ss_df[pub_ss_df.ImageWidth==IMAGE_SIZES[1]]\n    predict_df_3072 = pub_ss_df[pub_ss_df.ImageWidth==IMAGE_SIZES[2]]\n    predict_df_4096 = pub_ss_df[pub_ss_df.ImageWidth==IMAGE_SIZES[3]]\nelse:\n    # Load Segmentator\n    segmentator = cellsegmentator.CellSegmentator(NUC_MODEL, CELL_MODEL, scale_factor=0.25, padding=True)\n    \n    # Make subset dataframes\n    predict_df_1728 = ss_df[ss_df.ImageWidth==IMAGE_SIZES[0]]\n    predict_df_2048 = ss_df[ss_df.ImageWidth==IMAGE_SIZES[1]]\n    predict_df_3072 = ss_df[ss_df.ImageWidth==IMAGE_SIZES[2]]\n    predict_df_4096 = ss_df[ss_df.ImageWidth==IMAGE_SIZES[3]]\n\n\npredict_ids_1728 = predict_df_1728.ID.to_list()\npredict_ids_2048 = predict_df_2048.ID.to_list()\npredict_ids_3072 = predict_df_3072.ID.to_list()\npredict_ids_4096 = predict_df_4096.ID.to_list()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.600629Z","iopub.execute_input":"2025-09-25T06:58:56.600844Z","iopub.status.idle":"2025-09-25T06:58:56.62002Z","shell.execute_reply.started":"2025-09-25T06:58:56.600826Z","shell.execute_reply":"2025-09-25T06:58:56.619166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)  # 如果使用多个GPU","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.621052Z","iopub.execute_input":"2025-09-25T06:58:56.621361Z","iopub.status.idle":"2025-09-25T06:58:56.633212Z","shell.execute_reply.started":"2025-09-25T06:58:56.621326Z","shell.execute_reply":"2025-09-25T06:58:56.632553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TILE_SIZE = (256,256)\nDO_TTA = True\nTTA_Times = 6\nseed_everything(443)\n# model = DML_MIL2(mode='test',backbone='resnet50d')\n# model = LC_MIL2(backbone='resnet50d',layers=6,mode='test') ######\n# model = Cellnet(model_name='resnet50d', in_chan_num=4, class_num=1) # negative model\n# model =  timm.create_model(model_name='resnet50d', num_classes=19,\n#                               in_chans=4)\n\n# model = timm.create_model(model_name='convnextv2_base.fcmae_ft_in22k_in1k_384', num_classes=19, in_chans=4)\n# model = ResNestXdual()\n# model = MILResNet50D(model_name='resnet50d')\n\n# model = Chanel_MIL(layers=(4,4))\n\n# 去掉不匹配的前缀m , ll\n# model = JakiroResNet200D(model_name='resnet50d')\n# state_dict = torch.load('D:/DevProj/Singlecell_Classification/hpa_results/resnet50d/checkpoints/f1_epoch-15.pth')\n# state_dict = torch.load('D:/DevProj/Singlecell_Classification/hpa_results/cls0707/model_e16.pth')\n\n# state_dict = torch.load('/kaggle/input/lc-sim/model_e30.pth')\n# new_state_dict = collections.OrderedDict()\n# for k, v in state_dict.items():\n#     new_key = k.replace('_orig_mod.', '')\n#     new_state_dict[new_key] = v\n# model.load_state_dict(new_state_dict)\n# model = model.cuda()\n\n# model = Encoder(\n#         backbone=Backbone(\n#             \"resnet50d\",\n#             4,\n#         ),\n#         d_model=512,\n#         out_features=19,\n#         dropout=0.5,\n#         pos_dim=256,\n#         transformer_layers=6,\n#         transformer_heads=8,\n#         proj_dim=256,\n#         hid_dim=512,\n#     )\n# state_dict = torch.load(\"/kaggle/input/supcon-v3/model_25.pth\")\n# new_state_dict = collections.OrderedDict()\n# prefix = \"module._orig_mod.encoder_q.\"\n# for k, v in state_dict.items():\n#     if k.startswith(prefix):\n#         new_key = k[len(prefix) :]  # 去掉前缀，保留子模块结构\n#         new_state_dict[new_key] = v\n# model.load_state_dict(new_state_dict)\n# model = model.cuda()\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nMODEL_PATH = \"/kaggle/input/help-base/help_e20.pt\"\nmodel = torch.jit.load(MODEL_PATH,map_location=device)\nmodel = model.cuda()\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:56.634069Z","iopub.execute_input":"2025-09-25T06:58:56.634796Z","iopub.status.idle":"2025-09-25T06:58:59.856218Z","shell.execute_reply.started":"2025-09-25T06:58:56.634776Z","shell.execute_reply":"2025-09-25T06:58:59.855382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.transforms.v2 as v2\nmodel.eval()\nmodel_level = 'cell' #  image or cell \npredictions = []\nsub_df = pd.DataFrame(columns=[\"ID\"], data=predict_ids_1728+predict_ids_2048+predict_ids_3072+predict_ids_4096)\n\n# #### STEP TIMING FOR 1728x1728 IMAGES FOR EFFNETB0 ON 128x128 CROPS ####\n#  0:\t 1.03042 seconds\n#  1:\t 8.14935 seconds\n#  2:\t 0.00002 seconds\n#  3:\t 29.9057 seconds\n#  4:\t 1.30675 seconds\n#  5:\t 0.01442 seconds\n#  6:\t 0.26723 seconds\n#  7:\t 4.10871 seconds\n#  8:\t 0.00108 seconds\n#  9:\t 0.00066 seconds\n# 10:\t 0.00015 seconds\nto_tensor = transforms.Compose([\n            ToTensor(),\n            Normalize(mean=[0.485, 0.456, 0.406, 0.406], std=[0.229, 0.224, 0.225, 0.225]),\n        ])\n\nclass MinMaxNormalize(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, x):\n        min_val = torch.amin(x, dim=(1, 2), keepdim=True)\n        max_val = torch.amax(x, dim=(1, 2), keepdim=True)\n        return (x - min_val) / (max_val - min_val + 1e-6)\n        \ntransforms_pack = [\n    v2.RandomAffine(\n        degrees=(-180, 180), \n        translate=None,         # 禁用平移\n        scale=(0.9, 1.1),       # 随机缩放\n        shear=15                # x/y 轴剪切范围为 ±15 度\n    ),\n    v2.RandomHorizontalFlip(p=0.4),  # 40% 水平翻转概率\n    v2.RandomVerticalFlip(p=0.4),    # 40% 垂直翻转概率\n]\n\ntta_transforms = v2.Compose([\n    v2.ToImage(),\n    v2.ToDtype(torch.float32, scale=True),  # 缩放像素值到 [0, 1]\n    v2.RandomOrder(transforms_pack),  # 随机顺序执行增强\n    # v2.Normalize(\n    #     mean=[0.485, 0.456, 0.406, 0.406],  # 四通道均值\n    #     std=[0.229, 0.224, 0.225, 0.225],   # 四通道标准差\n    # )\n    v2.Normalize(mean=[0.08069, 0.05258, 0.05487, 0.08282],\n                     std=[0.13704, 0.10145, 0.15313, 0.13814])\n    # v2.Lambda(MinMaxNormalize())\n])\nfor index, row in tqdm(sub_df.iterrows(), total=len(sub_df)):\n    submission_id = row[\"ID\"]\n    image_width = pub_ss_df.loc[pub_ss_df['ID']==submission_id,'ImageWidth'].values[0]\n    image_height = pub_ss_df.loc[pub_ss_df['ID']==submission_id,'ImageHeight'].values[0]\n\n    # Step 0: Get batch of images as numpy arrays\n    images = load_img(TEST_IMG_DIR, submission_id) #to-do\n\n    if ONLY_PUBLIC:\n        # Step 1: Get Bounding Boxes\n        cell_bboxes = pub_ss_df.loc[pub_ss_df.ID==submission_id, 'mask_bboxes'].values\n\n        # Step 3: Get Submission RLEs\n        submission_rles = pub_ss_df.loc[pub_ss_df.ID==submission_id, 'mask_sub_rles'].values\n        # Optional Step: Get the Masks\n#         if IS_DEMO:\n        mask_rles = pub_ss_df.loc[pub_ss_df.ID==submission_id,'mask_rles'].values[0]\n        # mask = sum([rle_to_mask(mask, image_width, image_height) for mask in mask_rles ])\n        mask = np.zeros((image_height, image_width), dtype=np.uint16)\n\n        # 遍历每个 RLE 编码，赋予独立标签\n        for idx, rle_str in enumerate(mask_rles, start=1):\n            # 生成二值掩码（0 和 255）\n            binary_mask = rle_to_mask(rle_str, image_width, image_height)\n            # 转换为布尔掩码（True/False），并直接标记为当前标签值\n            mask[binary_mask > 0] = idx  # idx从1开始，背景保持为0\n        # plt.imshow(mask)\n        # plt.show()\n        # break\n\n    else:\n        # Step 1: Do Prediction On Batch\n        cell_segmentations = segmentator.pred_cells([[rgby_image[j] for rgby_image in batch_rgby_images] for j in [0, 3, 2]])\n        nuc_segmentations = segmentator.pred_nuclei([rgby_image[2] for rgby_image in batch_rgby_images])\n\n        # Step 2: Perform Cell Labelling on Batch\n        batch_masks = [label_cell(nuc_seg, cell_seg)[1].astype(np.uint8) for nuc_seg, cell_seg in zip(nuc_segmentations, cell_segmentations)]\n\n        # Step 3: Reshape the RGBY Images so They Are Channels Last Across the Batch\n        batch_rgb_images = [rgby_image.transpose(1,2,0)[..., :-1] for rgby_image in batch_rgby_images]\n\n        # Step 4: Get Bounding Boxes For All Cells in All Images in Batch\n        batch_cell_bboxes = [get_contour_bbox_from_raw(mask) for mask in batch_masks]\n\n        # Step 5: Generate Submission RLEs For the Batch\n        submission_rles = [[binary_mask_to_ascii(mask, mask_val=cell_id) for cell_id in range(1, mask.max()+1)] for mask in batch_masks]\n\n    \n    # Step 6: Cut Out, Pad to Square, and Resize to 224x224\n    # cell_tiles = [\n    #     cv2.resize(\n    #         pad_to_square(\n    #             images[bbox[1]:bbox[3], bbox[0]:bbox[2], ...]),\n    #         TILE_SIZE, interpolation=cv2.INTER_CUBIC) for bbox in cell_bboxes[0]]\n    cell_tiles = split_cells(images,mask,TILE_SIZE)\n    # print(len(cell_tiles))\n    # break\n\n    # Step 7: (OPTIONAL) Test Time Augmentation\n    if DO_TTA:\n        # List to store augmented images and predictions\n        tta_preds = []\n        # Apply augmentations and perform inference\n        for _ in range(TTA_Times):\n            augmented_images = [tta_transforms(cell_img) for cell_img in cell_tiles]\n            augmented_images = torch.stack(augmented_images, dim=0).cuda()  # [cell_cnt, 4, 224, 224]\n            \n            with torch.no_grad():\n                # with torch.cuda.amp.autocast():\n                if model_level == 'cell':\n                    cell_logits, feat = model(augmented_images)\n                    _preds = torch.sigmoid(cell_logits)\n                    tta_preds.append(_preds.cpu().detach().numpy())\n                else:\n                    cell_logits, query_logits = model(augmented_images, len(augmented_images))\n                    cell_preds = torch.sigmoid(cell_logits)\n                    query_preds = torch.sigmoid(query_logits)\n                    preds = cell_preds * query_preds  # [cell_cnt,19]\n                    preds[:, 11] = cell_preds[:, 11]\n                    preds[:, 18] = cell_preds[:, 18]\n                    preds = preds.cpu().detach().numpy()\n                    tta_preds.append(preds)\n    \n        # Average predictions from all augmentations\n        preds = np.mean(tta_preds, axis=0)\n    else:\n        cell_images = [to_tensor(cell_img) for cell_img in cell_tiles]\n    \n        # Step 8: Perform Inference\n        cell_cnt = len(cell_images)\n        cell_images = torch.stack(cell_images, dim=0).cuda()  # [cell_cnt, 4, 224, 224]\n        if model_level == 'cell':\n            preds = np.zeros((cell_cnt, 19))\n            with torch.no_grad():\n                with torch.cuda.amp.autocast():\n                    cell_logits = model(cell_images)\n            _preds = torch.sigmoid(cell_logits)\n            preds = _preds.cpu().detach().numpy()\n        else:\n            with torch.no_grad():\n                with torch.cuda.amp.autocast():\n                    cell_logits, query_logits = model(cell_images, cell_cnt)\n            cell_preds = torch.sigmoid(cell_logits)\n            query_preds = torch.sigmoid(query_logits)\n            preds = cell_preds * query_preds  # [cell_cnt,19]\n            preds[:, 11] = cell_preds[:, 11]\n            preds[:, 18] = cell_preds[:, 18]\n            preds = preds.cpu().detach().numpy()\n\n\n    pred_string = []\n    for cell_pred, mask_rle in zip(preds, submission_rles[0]):\n        for i in range(len(cell_pred)):\n            conf = cell_pred[i]\n            # if conf<CONF_THRESH:\n            #     continue\n            pred_string.append(f\"{i} {conf} {mask_rle}\")\n        # pred_string.append(f\"{18} {cell_pred[18]} {mask_rle}\")\n    pred_string = \" \".join(pred_string)\n    sub_df.loc[sub_df[\"ID\"]==submission_id, \"PredictionString\"] = pred_string # update one image's prediction string\n\nprint(\"\\n... TEST DATAFRAME ...\\n\")\ndisplay(sub_df.head(3))","metadata":{"papermill":{"duration":0.033138,"end_time":"2021-02-02T02:51:44.910682","exception":false,"start_time":"2021-02-02T02:51:44.877544","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:58:59.857213Z","iopub.execute_input":"2025-09-25T06:58:59.857437Z","iopub.status.idle":"2025-09-25T06:59:01.342529Z","shell.execute_reply.started":"2025-09-25T06:58:59.857421Z","shell.execute_reply":"2025-09-25T06:59:01.341431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ss_df = ss_df.merge(sub_df, how=\"left\", on=\"ID\")\nss_df[\"PredictionString\"] = ss_df.apply(create_pred_col, axis=1)\nss_df = ss_df.drop(columns=[\"PredictionString_x\", \"PredictionString_y\"])\nss_df.to_csv(\"/kaggle/working/submission.csv\", index=False)\ndisplay(ss_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-25T06:59:01.343236Z","iopub.status.idle":"2025-09-25T06:59:01.343556Z","shell.execute_reply.started":"2025-09-25T06:59:01.343379Z","shell.execute_reply":"2025-09-25T06:59:01.34339Z"}},"outputs":[],"execution_count":null}]}