{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10691327,"sourceType":"datasetVersion","datasetId":6624465},{"sourceId":11584716,"sourceType":"datasetVersion","datasetId":7263650}],"dockerImageVersionId":30805,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CZII inference___SPBU","metadata":{}},{"cell_type":"markdown","source":"Project was created by Kong Qi, Liu Jie, Xiong Xiao and Fan Xuanping and scientific adviser Petrosian Ovanes from Saint Petersburg State University, Russia.\n\nhttps://www.kaggle.com/liujiecool\nhttps://www.kaggle.com/kongqi123456\nhttps://www.kaggle.com/xiongxiao1\nhttps://www.kaggle.com/work\nhttps://www.kaggle.com/ovanespetrosian\n\n1. our score:.0.77226, top1:0.78759\n   \n2. There were 900 people on leadboard. The place of our project in this competetion is TOP9, it is the first 1% in the competetion.","metadata":{}},{"cell_type":"markdown","source":"## import, setting","metadata":{}},{"cell_type":"code","source":"!curl -v https://google.com","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:06:49.191963Z","iopub.execute_input":"2025-04-27T16:06:49.192323Z","iopub.status.idle":"2025-04-27T16:06:50.269381Z","shell.execute_reply.started":"2025-04-27T16:06:49.192287Z","shell.execute_reply":"2025-04-27T16:06:50.268572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!curl -v https://pypi.org","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:06:50.271013Z","iopub.execute_input":"2025-04-27T16:06:50.271297Z","iopub.status.idle":"2025-04-27T16:06:51.310082Z","shell.execute_reply.started":"2025-04-27T16:06:50.271268Z","shell.execute_reply":"2025-04-27T16:06:51.309275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install connected-components-3d ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:06:51.311394Z","iopub.execute_input":"2025-04-27T16:06:51.311700Z","iopub.status.idle":"2025-04-27T16:07:01.637995Z","shell.execute_reply.started":"2025-04-27T16:06:51.311670Z","shell.execute_reply":"2025-04-27T16:07:01.636655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q zarr ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:07:01.640989Z","iopub.execute_input":"2025-04-27T16:07:01.641397Z","iopub.status.idle":"2025-04-27T16:07:12.592966Z","shell.execute_reply.started":"2025-04-27T16:07:01.641350Z","shell.execute_reply":"2025-04-27T16:07:12.591731Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q copick[all]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:07:12.594693Z","iopub.execute_input":"2025-04-27T16:07:12.595486Z","iopub.status.idle":"2025-04-27T16:07:51.149784Z","shell.execute_reply.started":"2025-04-27T16:07:12.595440Z","shell.execute_reply":"2025-04-27T16:07:51.148961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q starfile ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:07:51.151004Z","iopub.execute_input":"2025-04-27T16:07:51.151261Z","iopub.status.idle":"2025-04-27T16:07:59.772712Z","shell.execute_reply.started":"2025-04-27T16:07:51.151234Z","shell.execute_reply":"2025-04-27T16:07:59.771501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q mrcfile>=1.4.3 ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:07:59.774608Z","iopub.execute_input":"2025-04-27T16:07:59.775275Z","iopub.status.idle":"2025-04-27T16:08:08.296357Z","shell.execute_reply.started":"2025-04-27T16:07:59.775197Z","shell.execute_reply":"2025-04-27T16:08:08.295430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport yaml\nimport sys\nimport cv2\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom glob import glob\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.optim import AdamW, Adam\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport pandas.api.types\nimport sklearn.metrics\nimport timm\nimport scipy\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\nimport json\nimport zarr, copick\nimport cc3d\nimport gc\nfrom joblib import Parallel, delayed\nfrom sklearn.cluster import DBSCAN","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:08.297882Z","iopub.execute_input":"2025-04-27T16:08:08.298268Z","iopub.status.idle":"2025-04-27T16:08:21.583328Z","shell.execute_reply.started":"2025-04-27T16:08:08.298224Z","shell.execute_reply":"2025-04-27T16:08:21.582621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 126 # friend's birthday\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndef seed_everything(seed):\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.backends.cudnn.deterministic = True # Fix the network according to random seed\n    print('Finish seeding with seed {}'.format(seed))\n\nseed_everything(SEED)\nprint('Training on device {}'.format(device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.584443Z","iopub.execute_input":"2025-04-27T16:08:21.585337Z","iopub.status.idle":"2025-04-27T16:08:21.631179Z","shell.execute_reply.started":"2025-04-27T16:08:21.585295Z","shell.execute_reply":"2025-04-27T16:08:21.630202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile config.yaml\n\ndebug: False\ndata_path: \"/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/\"\nbs: 2\nprogress_bar_refresh_rate: 1\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.635812Z","iopub.execute_input":"2025-04-27T16:08:21.636214Z","iopub.status.idle":"2025-04-27T16:08:21.673990Z","shell.execute_reply.started":"2025-04-27T16:08:21.636165Z","shell.execute_reply":"2025-04-27T16:08:21.672944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"config.yaml\", \"r\") as file_obj:\n    config = yaml.safe_load(file_obj)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.675194Z","iopub.execute_input":"2025-04-27T16:08:21.675552Z","iopub.status.idle":"2025-04-27T16:08:21.685305Z","shell.execute_reply.started":"2025-04-27T16:08:21.675514Z","shell.execute_reply":"2025-04-27T16:08:21.684079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_D = 32\n#256\nPATCH_H = 320\n#256\nPATCH_W = 320","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.686692Z","iopub.execute_input":"2025-04-27T16:08:21.687012Z","iopub.status.idle":"2025-04-27T16:08:21.696076Z","shell.execute_reply.started":"2025-04-27T16:08:21.686976Z","shell.execute_reply":"2025-04-27T16:08:21.694960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"COOR_LIST = []\nfor x in range(630//PATCH_W+1):\n    for y in range(630//PATCH_H+1):\n        for z in range(184//PATCH_D+2):\n            COOR_LIST.append([x*PATCH_W-10*x, y*PATCH_H-10*y, z*PATCH_D-6*z])\n            ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.697271Z","iopub.execute_input":"2025-04-27T16:08:21.697981Z","iopub.status.idle":"2025-04-27T16:08:21.706207Z","shell.execute_reply.started":"2025-04-27T16:08:21.697943Z","shell.execute_reply":"2025-04-27T16:08:21.705471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass CZIIDataset(Dataset):\n    def __init__(self, data_path, obj, usage='sub'):\n        self.obj = obj\n        self.data_path = data_path\n        self.usage = usage\n        self.coor_list = COOR_LIST\n        self.voxel = np.zeros((184, 630, 630), dtype=np.float32)\n\n        # loading and normalization\n        self.voxel[:, :, :] = self.normalize_numpy(\n            zarr.open(data_path + obj + '/VoxelSpacing10.000/denoised.zarr')['0'][:]\n        )\n        \n    def __getitem__(self, index):\n        row = self.coor_list[index]\n        x = row[0]\n        y = row[1]\n        z = row[2]\n\n        if z > 184-PATCH_D:\n            z = 184-PATCH_D\n        if x > 630-PATCH_W:\n            x = 630-PATCH_W\n        if y > 630-PATCH_H:\n            y = 630-PATCH_H\n        data = self.voxel[z:z + PATCH_D, y:y+PATCH_H, x:x+PATCH_W]\n        data = torch.tensor(data, dtype=torch.float32)\n\n        return data.unsqueeze(0), torch.tensor(x), torch.tensor(y), torch.tensor(z)\n\n    def normalize_numpy(self, x):\n        lower, upper = np.percentile(x, (1, 99))\n        x = np.clip(x, lower, upper)\n        x = x - np.min(x)\n        x = x / np.max(x)\n        return x\n\n    def __len__(self):\n        return len(self.coor_list)","metadata":{"execution":{"iopub.status.busy":"2025-04-27T16:08:21.707023Z","iopub.execute_input":"2025-04-27T16:08:21.707258Z","iopub.status.idle":"2025-04-27T16:08:21.718326Z","shell.execute_reply.started":"2025-04-27T16:08:21.707235Z","shell.execute_reply":"2025-04-27T16:08:21.717534Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## eval","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\"\"\"\nDerived from:\nhttps://github.com/cellcanvas/album-catalog/blob/main/solutions/copick/compare-picks/solution.py\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\n\nfrom scipy.spatial import KDTree\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef compute_metrics(reference_points, reference_radius, candidate_points):\n    num_reference_particles = len(reference_points)\n    num_candidate_particles = len(candidate_points)\n\n    if len(reference_points) == 0:\n        return 0, num_candidate_particles, 0\n\n    if len(candidate_points) == 0:\n        return 0, 0, num_reference_particles\n\n    ref_tree = KDTree(reference_points)\n    candidate_tree = KDTree(candidate_points)\n    raw_matches = candidate_tree.query_ball_tree(ref_tree, r=reference_radius)\n    matches_within_threshold = []\n    for match in raw_matches:\n        matches_within_threshold.extend(match)\n    # Prevent submitting multiple matches per particle.\n    # This won't be be strictly correct in the (extremely rare) case where true particles\n    # are very close to each other.\n    matches_within_threshold = set(matches_within_threshold)\n    tp = int(len(matches_within_threshold))\n    fp = int(num_candidate_particles - tp)\n    fn = int(num_reference_particles - tp)\n    return tp, fp, fn\n\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        distance_multiplier: float,\n        beta: int) -> float:\n    '''\n    F_beta\n      - a true positive occurs when\n         - (a) the predicted location is within a threshold of the particle radius, and\n         - (b) the correct `particle_type` is specified\n      - raw results (TP, FP, FN) are aggregated across all experiments for each particle type\n      - f_beta is calculated for each particle type\n      - individual f_beta scores are weighted by particle type for final score\n    '''\n\n    particle_radius = {\n        'apo-ferritin': 60,\n        'beta-amylase': 65,\n        'beta-galactosidase': 90,\n        'ribosome': 150,\n        'thyroglobulin': 130,\n        'virus-like-particle': 135,\n    }\n\n    weights = {\n        'apo-ferritin': 1,\n        'beta-amylase': 0,\n        'beta-galactosidase': 2,\n        'ribosome': 1,\n        'thyroglobulin': 2,\n        'virus-like-particle': 1,\n    }\n\n    particle_radius = {k: v * distance_multiplier for k, v in particle_radius.items()}\n\n    # Filter submission to only contain experiments found in the solution split\n    split_experiments = set(solution['experiment'].unique())\n    submission = submission.loc[submission['experiment'].isin(split_experiments)]\n\n    # Only allow known particle types\n    if not set(submission['particle_type'].unique()).issubset(set(weights.keys())):\n        raise ParticipantVisibleError('Unrecognized `particle_type`.')\n\n    assert solution.duplicated(subset=['experiment', 'x', 'y', 'z']).sum() == 0\n    assert particle_radius.keys() == weights.keys()\n\n    results = {}\n    for particle_type in solution['particle_type'].unique():\n        results[particle_type] = {\n            'total_tp': 0,\n            'total_fp': 0,\n            'total_fn': 0,\n        }\n\n    for experiment in split_experiments:\n        for particle_type in solution['particle_type'].unique():\n            reference_radius = particle_radius[particle_type]\n            select = (solution['experiment'] == experiment) & (solution['particle_type'] == particle_type)\n            reference_points = solution.loc[select, ['x', 'y', 'z']].values\n\n            select = (submission['experiment'] == experiment) & (submission['particle_type'] == particle_type)\n            candidate_points = submission.loc[select, ['x', 'y', 'z']].values\n\n            if len(reference_points) == 0:\n                reference_points = np.array([])\n                reference_radius = 1\n\n            if len(candidate_points) == 0:\n                candidate_points = np.array([])\n\n            tp, fp, fn = compute_metrics(reference_points, reference_radius, candidate_points)\n\n            results[particle_type]['total_tp'] += tp\n            results[particle_type]['total_fp'] += fp\n            results[particle_type]['total_fn'] += fn\n\n    aggregate_fbeta = 0.0\n    partial_fbeta = {}\n    for particle_type, totals in results.items():\n        tp = totals['total_tp']\n        fp = totals['total_fp']\n        fn = totals['total_fn']\n\n        precision = tp / (tp + fp) if tp + fp > 0 else 0\n        recall = tp / (tp + fn) if tp + fn > 0 else 0\n        fbeta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall) if (precision + recall) > 0 else 0.0\n        aggregate_fbeta += fbeta * weights.get(particle_type, 1.0)\n        partial_fbeta[particle_type] = fbeta\n\n    if weights:\n        aggregate_fbeta = aggregate_fbeta / sum(weights.values())\n    else:\n        aggregate_fbeta = aggregate_fbeta / len(results)\n    return aggregate_fbeta, partial_fbeta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.719566Z","iopub.execute_input":"2025-04-27T16:08:21.719879Z","iopub.status.idle":"2025-04-27T16:08:21.737736Z","shell.execute_reply.started":"2025-04-27T16:08:21.719854Z","shell.execute_reply":"2025-04-27T16:08:21.736999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Models","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass ConvNeXtBlock3D(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dwconvx3 = nn.Conv3d(dim, dim, kernel_size=3, padding=1, groups=dim, bias=False)\n        self.norm = nn.GroupNorm(num_groups=1, num_channels=dim)\n        self.pwconv1 = nn.Conv3d(dim, dim*4, kernel_size=1, bias=False)\n        self.act = nn.GELU()\n        self.pwconv2 = nn.Conv3d(dim*4, dim, kernel_size=1, bias=False)\n\n    def forward(self, x):\n        input = x\n        x = self.dwconvx3(x)\n        x = self.norm(x)\n        x = self.pwconv1(x)\n        x = self.act(x)\n        x = self.pwconv2(x)\n        x = input + x\n        return x\n\nclass Stem3D(nn.Module):\n    def __init__(self, in_channels, out_channels, drop_path=0.0, layer_scale_init_value=1e-6):\n        super().__init__()\n        self.stem = nn.Conv3d(in_channels, out_channels, kernel_size=2, stride=2)\n        self._norm = nn.GroupNorm(num_groups=1, num_channels=out_channels)\n\n    def forward(self, x):\n        x = self.stem(x)\n        x = self._norm(x)\n        return x\n\nclass ConvNeXt3DEncoderUNetStyle(nn.Module):\n    def __init__(self,\n                 in_ch=1,\n                 dims=[64, 128, 256, 512],\n                 depths=[3, 3, 3, 3]):\n        super().__init__()\n        # Stem\n        self.stem = Stem3D(in_ch, dims[0])\n\n        self.downsample_layers = nn.ModuleList([\n            nn.Identity(), # no downsampling after stem\n            nn.Sequential(\n                nn.GroupNorm(num_groups=1, num_channels=dims[0]),\n                nn.Conv3d(dims[0], dims[1], kernel_size=(2, 2, 2), stride=(2, 2, 2)),\n            ),\n            nn.Sequential(\n                nn.GroupNorm(num_groups=1, num_channels=dims[1]),\n                nn.Conv3d(dims[1], dims[2], kernel_size=(2, 2, 2), stride=(2, 2, 2)),\n            ),\n            nn.Sequential(\n                nn.GroupNorm(num_groups=1, num_channels=dims[2]),\n                nn.Conv3d(dims[2], dims[3], kernel_size=(2, 2, 2), stride=(2, 2, 2)),\n            )\n        ])\n\n        cur = 0\n        self.stages = nn.ModuleList()\n        for i in range(4):\n            blocks = []\n            for j in range(depths[i]):\n                blocks.append(ConvNeXtBlock3D(dim=dims[i]))\n            cur += depths[i]\n            self.stages.append(nn.Sequential(*blocks))\n        self.norm = nn.GroupNorm(num_groups=1, num_channels=dims[-1])\n\n    def forward(self, x):\n        f0 = self.stem(x)\n\n        # stage0\n        # no downsampling\n        f0 = self.stages[0](f0)\n        # stage1\n        f1 = self.downsample_layers[1](f0)\n        f1 = self.stages[1](f1)\n        # stage2\n        f2 = self.downsample_layers[2](f1)\n        f2 = self.stages[2](f2)\n        # stage3\n        f3 = self.downsample_layers[3](f2)\n        f3 = self.stages[3](f3)\n        f3 = self.norm(f3)\n\n        return [x, f0, f1, f2, f3]\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass ConvBlock3D(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_channels, in_channels, kernel_size=3, padding=1, bias=False, groups=in_channels)\n        self.norm1 = nn.GroupNorm(num_groups=1, num_channels=in_channels)\n        self.act1 = nn.GELU()\n        self.conv2 = nn.Conv3d(in_channels, in_channels*3, kernel_size=1, bias=False)\n        self.conv3 = nn.Conv3d(in_channels*3, out_channels, kernel_size=1, bias=False)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.norm1(x)\n        x = self.conv2(x)\n        x = self.act1(x)\n        x = self.conv3(x)\n        return x\n\n\nclass Decoder3D(nn.Module):\n    def __init__(self,\n                 encoder_dims=[64, 128, 256, 512],\n                 decoder_dims=[32, 64, 128, 256],\n                 out_channels=1):\n        super().__init__()\n\n        # 1) f3 -> f2\n        self.up3 = nn.Upsample(scale_factor=(2,2,2), mode='trilinear', align_corners=True)\n        self.dec3 = ConvBlock3D(in_channels=encoder_dims[3] + encoder_dims[2], out_channels=decoder_dims[2])\n        # 2) f2 -> f1\n        self.up2 = nn.Upsample(scale_factor=(2,2,2), mode='trilinear', align_corners=True)\n        self.dec2 = ConvBlock3D(in_channels=decoder_dims[2] + encoder_dims[1], out_channels=decoder_dims[1])\n        # 3) f1 -> f0\n        self.up1 = nn.Upsample(scale_factor=(2,2,2), mode='trilinear', align_corners=True)\n        self.dec1 = ConvBlock3D(in_channels=decoder_dims[1] + encoder_dims[0], out_channels=decoder_dims[0])\n        # 4) f0 -> x\n        self.up0 = nn.Upsample(scale_factor=(2,2,2), mode='trilinear', align_corners=True)\n        self.dec0 = ConvBlock3D(in_channels=decoder_dims[0], out_channels=decoder_dims[0])\n        self.final_conv = nn.Conv3d(decoder_dims[0], out_channels, kernel_size=1)\n\n    def forward(self, features):\n        x, f0, f1, f2, f3 = features\n\n        # --- 1) f3 -> f2 ---\n        d3 = self.up3(f3)                  # Upsample\n        d3 = torch.cat([d3, f2], dim=1)    # skip connect\n        d3 = self.dec3(d3)                 # conv block\n\n        # --- 2) f2 -> f1 ---\n        d2 = self.up2(d3)\n        d2 = torch.cat([d2, f1], dim=1)\n        d2 = self.dec2(d2)\n\n        # --- 3) f1 -> f0 ---\n        d1 = self.up1(d2)\n        d1 = torch.cat([d1, f0], dim=1)\n        d1 = self.dec1(d1)\n\n        # --- 4) f0 -> x ---\n        d0 = self.up0(d1)\n        d0 = self.dec0(d0)\n\n        out = self.final_conv(d0)\n        return out\n\nclass UNet3D(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.encoder = ConvNeXt3DEncoderUNetStyle(in_channels)\n        self.decoder = Decoder3D(out_channels=out_channels)\n    def forward(self, x):\n        features = self.encoder(x)\n        out = self.decoder(features)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.738946Z","iopub.execute_input":"2025-04-27T16:08:21.739171Z","iopub.status.idle":"2025-04-27T16:08:21.761372Z","shell.execute_reply.started":"2025-04-27T16:08:21.739148Z","shell.execute_reply":"2025-04-27T16:08:21.760758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CZIIModule(pl.LightningModule):\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        self.model = UNet3D(in_channels=1, out_channels=5)\n\n    def forward(self, batch):\n        preds = self.model(batch)\n        return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.762363Z","iopub.execute_input":"2025-04-27T16:08:21.762611Z","iopub.status.idle":"2025-04-27T16:08:21.774562Z","shell.execute_reply.started":"2025-04-27T16:08:21.762577Z","shell.execute_reply":"2025-04-27T16:08:21.773882Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"import numpy as np \na = np.array([0.5, 0.1, 0.05, 0.1, 0.3]) # TS_69_2 v16\nb = np.array([0.3, 0.1, 0.15, 0.05, 0.2]) # val2\nc = np.array([0.1, 0.05, 0.15, 0.2, 0.15]) # TS_69_2 best\nd = np.array([0.3, 0.2, 0.3, 0.15, 0.3]) # all v13\ne = np.array([0.25, 0.25, 0.4, 0.05, 0.25]) # TS_86_3 best\nf = np.array([0.25, 0.25, 0.1, 0.05, 0.3]) #TS_73_6 v13\n(a+b+c+d+e+f)/6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.775487Z","iopub.execute_input":"2025-04-27T16:08:21.775776Z","iopub.status.idle":"2025-04-27T16:08:21.785409Z","shell.execute_reply.started":"2025-04-27T16:08:21.775751Z","shell.execute_reply":"2025-04-27T16:08:21.784594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nparticles = ['apo-ferritin', 'beta-galactosidase', 'ribosome', 'thyroglobulin',  'virus-like-particle']\nthreshold = torch.tensor([0.28333333, 0.15833333, 0.19166667, 0.1       , 0.25      ], device=device).reshape(5,1,1,1)\nmodel_path_list = [\n    '/kaggle/input/czii-9th-final-models/convnext-convnext_decoder-2valid-window_32_256_256-pre_ds-r5n5n33n33n33-4tta-bce-TS_69_2-v16.ckpt', \n    '/kaggle/input/czii-9th-final-models/convnext-convnext_decoder-2valid-window_32_256_256-pre_ds-r4n25-4tta-bce-2.ckpt',\n    '/kaggle/input/czii-9th-final-models/convnext-convnext_decoder-2valid-window_32_256_256-pre_ds-r5n5n33n33n33-4tta-bce-TS_69_2.ckpt', \n    '/kaggle/input/czii-9th-final-models/convnext-window_32_320_320-pre_ds-r5n5n33n33n33-4tta-bce-30e-all-v13.ckpt',\n    '/kaggle/input/czii-9th-final-models/convnext-convnext_decoder-2valid-window_32_256_256-pre_ds-r5n5n33n33n33-4tta-bce-TS_86_3.ckpt', \n    '/kaggle/input/czii-9th-final-models/convnext-convnext_decoder-2valid-window_32_256_256-pre_ds-r5n5n33n33n33-4tta-bce-TS_73_6-v13.ckpt', \n]\nmodel_list = []\nfor model_path in model_path_list: \n    model = CZIIModule.load_from_checkpoint(model_path, config=config, strict=False)\n    model.half()\n    model.eval()\n    for param in model.parameters():\n        param.grad = None\n    model = torch.nn.DataParallel(model, device_ids=[0, 1])\n    model.to(device)\n    \n    model_list.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:21.786428Z","iopub.execute_input":"2025-04-27T16:08:21.786742Z","iopub.status.idle":"2025-04-27T16:08:26.764981Z","shell.execute_reply.started":"2025-04-27T16:08:21.786698Z","shell.execute_reply":"2025-04-27T16:08:26.763982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calc_centroid(pred_temp):\n    component = cc3d.connected_components(pred_temp.copy(), connectivity=6)\n    stats = cc3d.statistics(component)\n    zyx = stats['centroids'][1:]\n    return np.ascontiguousarray(zyx[:, ::-1])*10.012444\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:08:26.766323Z","iopub.execute_input":"2025-04-27T16:08:26.766734Z","iopub.status.idle":"2025-04-27T16:08:26.772861Z","shell.execute_reply.started":"2025-04-27T16:08:26.766692Z","shell.execute_reply":"2025-04-27T16:08:26.771876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nexp_list = [p.split('/')[-1] for p in glob(config['data_path'] + '*', recursive=True)]\npred_df_list = []\nnum_model = len(model_list)\nfor exp in exp_list: \n    dataset = CZIIDataset(config['data_path'], exp)\n    data_loader = DataLoader(\n        dataset,\n        batch_size=config[\"bs\"],\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True, \n        persistent_workers=True\n        )\n    pred_temp = torch.zeros((5, 184, 630, 630), device=device, dtype=torch.float16)\n    wrap_temp = torch.zeros((5, 184, 630, 630), device=device, dtype=torch.float16)\n    with torch.no_grad():\n        for _data, x, y, z in data_loader: \n            _data = _data.to(device, non_blocking=True).half()\n            shape = _data.shape\n            _data = torch.cat([_data,\n                                torch.rot90(_data, k=1, dims=(3, 4)),\n                                torch.rot90(_data, k=2, dims=(3, 4)),\n                                torch.rot90(_data, k=3, dims=(3, 4)), \n                              ], dim=0)\n            pred_mask_list = []\n            for model in model_list: \n                preds = model.forward(_data)\n                pred_mask_list.append(preds)\n            preds = torch.stack(pred_mask_list, dim=0)\n            preds = preds.reshape(num_model, 4, shape[0], 5, *shape[2:])\n            preds = torch.stack([preds[:, 0],\n                                torch.rot90(preds[:, 1], k=3, dims=(4, 5)),\n                                torch.rot90(preds[:, 2], k=2, dims=(4, 5)),\n                                torch.rot90(preds[:, 3], k=1, dims=(4, 5)), \n                                ], dim=1).mean(1).sigmoid().sum(0)\n            for _x, _y, _z, _pred in zip(x, y, z, preds):\n                pred_temp[:, _z:_z+PATCH_D, _y:_y+PATCH_H, _x:_x+PATCH_W].add_(_pred)\n                wrap_temp[:, _z:_z+PATCH_D, _y:_y+PATCH_H, _x:_x+PATCH_W].add_(num_model)\n    pred_temp = (pred_temp/wrap_temp)\n    pred_temp = (pred_temp > threshold).detach().cpu().numpy()\n    del dataset, data_loader\n    gc.collect()\n    centroid = []\n    centroid = Parallel(n_jobs=5)([delayed(calc_centroid)(p_pred) for p_pred in pred_temp])\n    \n    for c, p in zip(centroid, particles):\n        centroid_dict = {'experiment': exp, 'particle_type': p, 'x': c[:,0], 'y': c[:,1], 'z': c[:,2]}\n        pred_df_list.append(pd.DataFrame(centroid_dict))\n    del centroid, centroid_dict\n    gc.collect()\n    ","metadata":{"execution":{"iopub.status.busy":"2025-04-27T16:08:26.774092Z","iopub.execute_input":"2025-04-27T16:08:26.774428Z","iopub.status.idle":"2025-04-27T16:12:16.655306Z","shell.execute_reply.started":"2025-04-27T16:08:26.774389Z","shell.execute_reply":"2025-04-27T16:12:16.654408Z"},"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_df = pd.concat(pred_df_list)\npred_df['id'] = range(len(pred_df))\npred_df = pred_df.reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:16.656495Z","iopub.execute_input":"2025-04-27T16:12:16.656792Z","iopub.status.idle":"2025-04-27T16:12:16.667661Z","shell.execute_reply.started":"2025-04-27T16:12:16.656765Z","shell.execute_reply":"2025-04-27T16:12:16.666918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# reference: https://www.kaggle.com/code/linheshen/model-ensemble-dbscan\n\n# 假设sub已经给定，拼接DataFrame\ndf = pred_df.copy()\n# 粒子半径映射\nparticle_names = ['apo-ferritin', 'beta-amylase', 'beta-galactosidase', 'ribosome', 'thyroglobulin', 'virus-like-particle']\nparticle_radius = {\n    'apo-ferritin': 60,\n    'beta-amylase': 65,\n    'beta-galactosidase': 90,\n    'ribosome': 150,\n    'thyroglobulin': 130,\n    'virus-like-particle': 135,\n}\n\nfinal = []  # 用于存储最终的结果\nfor pidx, p in enumerate(particle_names):\n    # 筛选出该粒子类型的所有点\n    pdf = df[df['particle_type'] == p].reset_index(drop=True)\n    p_rad = particle_radius[p]\n    \n    # 根据 experiment 分组\n    grouped = pdf.groupby(['experiment'])\n    \n    for exp, group in grouped:\n        group = group.reset_index(drop=True)\n        \n        # 使用DBSCAN进行聚类\n        coords = group[['x', 'y', 'z']].values\n        db = DBSCAN(eps=p_rad*0.5, min_samples=2, metric='euclidean', algorithm='kd_tree').fit(coords)\n        labels = db.labels_\n        \n        # 将聚类结果添加到DataFrame中\n        group['cluster'] = labels\n        \n        # 对每个簇进行处理\n        for cluster_id in np.unique(labels):\n            if cluster_id == -1:\n                continue  # 跳过噪声点\n            \n            cluster_points = group[group['cluster'] == cluster_id]\n            \n            # 计算簇的中心（平均位置）\n            avg_x = cluster_points['x'].mean()\n            avg_y = cluster_points['y'].mean()\n            avg_z = cluster_points['z'].mean()\n            \n            # 更新簇内点的位置\n            group.loc[group['cluster'] == cluster_id, ['x', 'y', 'z']] = avg_x, avg_y, avg_z\n            group = group.drop_duplicates(subset=['x', 'y', 'z'])\n        # 将处理后的数据添加到 final 列表\n        final.append(group)\n\n# 合并处理后的数据\npred_df = pd.concat(final, ignore_index=True)\npred_df = pred_df.drop(columns=['cluster'])\n\n# 排序按 'experiment' 和 'particle_type' 两列\npred_df = pred_df.sort_values(by=['experiment', 'particle_type']).reset_index(drop=True)\n\n# 重新生成 'id' 列，从 1 开始\npred_df['id'] = np.arange(0, len(pred_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:16.668553Z","iopub.execute_input":"2025-04-27T16:12:16.668903Z","iopub.status.idle":"2025-04-27T16:12:16.742585Z","shell.execute_reply.started":"2025-04-27T16:12:16.668856Z","shell.execute_reply":"2025-04-27T16:12:16.742004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:16.743583Z","iopub.execute_input":"2025-04-27T16:12:16.743929Z","iopub.status.idle":"2025-04-27T16:12:17.779723Z","shell.execute_reply.started":"2025-04-27T16:12:16.743890Z","shell.execute_reply":"2025-04-27T16:12:17.778735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:17.781077Z","iopub.execute_input":"2025-04-27T16:12:17.781396Z","iopub.status.idle":"2025-04-27T16:12:17.795538Z","shell.execute_reply.started":"2025-04-27T16:12:17.781366Z","shell.execute_reply":"2025-04-27T16:12:17.794948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if config['debug']: \n    df_list = []\n    pickable_obj_path_list = ['/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns/' + f'{exp}/Picks/*.json' for exp in ['TS_5_4', 'TS_69_2', 'TS_6_4']]\n    pickable_obj_path_list\n    for pickable_obj_path in pickable_obj_path_list:\n        path_list = glob(pickable_obj_path)\n        for path in path_list:\n            picks = json.load(open(path))\n            pickable_object_name = picks['pickable_object_name']\n            run_name = picks['run_name']\n            points = picks['points']\n            point_dict = {'x': [], 'y': [], 'z': []}\n            for p in points:\n                point_dict['x'].append(p['location']['x'])\n                point_dict['y'].append(p['location']['y'])\n                point_dict['z'].append(p['location']['z'])\n            df = pd.DataFrame(point_dict)\n            df['experiment'] = run_name\n            df['particle_type'] = pickable_object_name\n            df_list.append(df)\n    df = pd.concat(df_list)\n    df = df[['experiment', 'particle_type', 'x', 'y', 'z']]\n    display(df.head())\n    df['id'] = range(len(df))\n    df = df.reset_index(drop=True)\n    pseudo_score = score(df, pred_df, row_id_column_name='id', distance_multiplier=.5, beta=4)\n    print(pseudo_score)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:17.796410Z","iopub.execute_input":"2025-04-27T16:12:17.796681Z","iopub.status.idle":"2025-04-27T16:12:17.804053Z","shell.execute_reply.started":"2025-04-27T16:12:17.796625Z","shell.execute_reply":"2025-04-27T16:12:17.803255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_df.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-27T16:12:17.805065Z","iopub.execute_input":"2025-04-27T16:12:17.805364Z","iopub.status.idle":"2025-04-27T16:12:17.818895Z","shell.execute_reply.started":"2025-04-27T16:12:17.805339Z","shell.execute_reply":"2025-04-27T16:12:17.818193Z"}},"outputs":[],"execution_count":null}]}