{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":84969,"databundleVersionId":10033515},{"sourceType":"datasetVersion","sourceId":10703616,"datasetId":6633179,"databundleVersionId":11051385},{"sourceType":"datasetVersion","sourceId":9862305,"datasetId":6052780,"databundleVersionId":10114338},{"sourceType":"datasetVersion","sourceId":9867543,"datasetId":6040935,"databundleVersionId":10120220}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Improved private LB 0.51 --> 0.71 with very basic things!!\nThe private LB was improved without any changes of model architecture from [ahsuna123's notebook](https://www.kaggle.com/code/ahsuna123/3d-u-net-training-only). Sincere thanks to @ahsuna123 !!\n<br>What I mainly did:\n- increase ``max_epoch`` 50 --> 999\n- set LR scheduler multi-step\n- overlap with gaussian weight\n- optimize augmentation composite (next cell)\n- blob radius * 0.5 and thresh = 25 \n- pre-train with all tomo type data except ``denoised`` then retrain the model with ``denoised``\n- TTA rotate90 x3\n\nAnd train/val split is:<br>\n- train : TS_5_4, TS_69_2, TS_73_6, TS_86_3, TS_99_9<br>\n- val   : TS_6_4<br>\n\n<br>\nTrain code:<br>\nhttps://github.com/kyotaro-horio/Kaggle_CZII","metadata":{}},{"cell_type":"markdown","source":"## augmentation composite\n```python\nrandom_transforms = Compose([\n    RandCropByPosNegLabeld(\n        keys=[\"image\", \"label\"],\n        label_key=\"label\",\n        spatial_size=patch_size, \n        pos=1,\n        neg=1,\n        num_samples=cfg.batch_size,  \n        image_key=\"image\",\n        image_threshold=0\n    ),\n    RandRotate90d(keys=[\"image\", \"label\"], prob=0.5, spatial_axes=[0, 2]),\n    RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=0), \n    Rand3DElasticd(\n        keys=[\"image\", \"label\"], prob=0.2,\n        sigma_range=(2, 4), magnitude_range=(1, 2),\n        mode=(\"bilinear\", \"nearest\"), rotate_range=(0, 0, 0)  \n    ),\n    RandGaussianSmoothd(\n        keys=[\"image\"], prob=0.5,  \n        sigma_x=(0.5, 1.5), sigma_y=(0.5, 1.5), sigma_z=(0.5, 1.5), \n    ), \n    RandStdShiftIntensityd(keys=[\"image\"], prob=0.5, factors=0.1),\n    RandAdjustContrastd(keys=[\"image\"], prob=0.7, gamma=[0.8, 1.2]), \n])\n```","metadata":{}},{"cell_type":"markdown","source":"# Inference Code","metadata":{}},{"cell_type":"code","source":"deps_path = '/kaggle/input/czii-cryoet-dependencies'\n! cp -r /kaggle/input/czii-cryoet-dependencies/asciitree-0.3.3/ asciitree-0.3.3/\n! pip wheel asciitree-0.3.3/asciitree-0.3.3/\n! pip install -q asciitree-0.3.3-py3-none-any.whl\n! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:39:38.675337Z","iopub.execute_input":"2026-05-10T03:39:38.675602Z","iopub.status.idle":"2026-05-10T03:41:15.688336Z","shell.execute_reply.started":"2026-05-10T03:39:38.675577Z","shell.execute_reply":"2026-05-10T03:41:15.687258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CLASS_IDS = [1,2,3,4,5,6]\nNUM_CLASSES = len(CLASS_IDS) + 1\nID_TO_NAME = {\n    1: \"apo-ferritin\", \n    2: \"beta-amylase\",\n    3: \"beta-galactosidase\", \n    4: \"ribosome\", \n    5: \"thyroglobulin\", \n    6: \"virus-like-particle\"\n}\n\n# --- for inference ---\nMODEL_FOLDER = '/kaggle/input/20250209-202421-3dunet-16-1e-3-999-96x96x96'\nFOLD = 2 \nPATCH_SIZE = [96, 96, 96] #zyx\nCERT_THRESH = [0.5, 0.95, 0.95, 0.95, 0.55, 0.95, 0.75]\nBLOB_THRESH = 25\nOVERLAP = [2,2,2] #zyx\n\nTTA_K_ROTATE = 3\nDO_TTA = True if TTA_K_ROTATE>0 else False\n\n# --- copick ---\nCOPICK_CONFIG_PATH = \"/kaggle/working/copick.config\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:15.690628Z","iopub.execute_input":"2026-05-10T03:41:15.690935Z","iopub.status.idle":"2026-05-10T03:41:15.697321Z","shell.execute_reply.started":"2026-05-10T03:41:15.690908Z","shell.execute_reply":"2026-05-10T03:41:15.696417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List, Tuple, Union\nimport os\nimport numpy as np\nfrom pathlib import Path\nimport torch\nimport torchinfo\nimport zarr, copick\nfrom tqdm import tqdm\nfrom monai.data import DataLoader, Dataset, CacheDataset, decollate_batch\nfrom monai.transforms import (\n    Compose, \n    EnsureChannelFirstd, \n    Orientationd,  \n    AsDiscrete,  \n    RandFlipd, \n    RandRotate90d, \n    NormalizeIntensityd,\n    RandCropByLabelClassesd,\n    Rotate90, \n)\nfrom monai.networks.nets import UNet\n\nimport json\nimport cc3d\nimport time\nfrom glob import glob\n\nimport sys\nsys.path.append('/kaggle/input/hengck-czii-cryo-et-01')\n\nfrom czii_helper import *\nfrom dataset import *\n\nfrom scipy.optimize import linear_sum_assignment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:15.698393Z","iopub.execute_input":"2026-05-10T03:41:15.698729Z","iopub.status.idle":"2026-05-10T03:41:52.313926Z","shell.execute_reply.started":"2026-05-10T03:41:15.698704Z","shell.execute_reply":"2026-05-10T03:41:52.313112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config_blob = \"\"\"{\n    \"name\": \"czii_cryoet_mlchallenge_2024\",\n    \"description\": \"2024 CZII CryoET ML Challenge training data.\",\n    \"version\": \"1.0.0\",\n\n    \"pickable_objects\": [\n        {\n            \"name\": \"apo-ferritin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"4V1W\",\n            \"label\": 1,\n            \"color\": [  0, 117, 220, 128],\n            \"radius\": 60,\n            \"map_threshold\": 0.0418\n        },\n        {\n            \"name\": \"beta-amylase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"1FA2\",\n            \"label\": 2,\n            \"color\": [153,  63,   0, 128],\n            \"radius\": 65,\n            \"map_threshold\": 0.035\n        },\n        {\n            \"name\": \"beta-galactosidase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6X1Q\",\n            \"label\": 3,\n            \"color\": [ 76,   0,  92, 128],\n            \"radius\": 90,\n            \"map_threshold\": 0.0578\n        },\n        {\n            \"name\": \"ribosome\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6EK0\",\n            \"label\": 4,\n            \"color\": [  0,  92,  49, 128],\n            \"radius\": 150,\n            \"map_threshold\": 0.0374\n        },\n        {\n            \"name\": \"thyroglobulin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6SCJ\",\n            \"label\": 5,\n            \"color\": [ 43, 206,  72, 128],\n            \"radius\": 130,\n            \"map_threshold\": 0.0278\n        },\n        {\n            \"name\": \"virus-like-particle\",\n            \"is_particle\": true,\n            \"label\": 6,\n            \"color\": [255, 204, 153, 128],\n            \"radius\": 135,\n            \"map_threshold\": 0.201\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": true\n    },\n\n    \"static_root\": \"/kaggle/input/czii-cryo-et-object-identification/test/static\"\n}\"\"\"\nwith open(COPICK_CONFIG_PATH, \"w\") as f:\n    f.write(config_blob)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:52.315270Z","iopub.execute_input":"2026-05-10T03:41:52.316262Z","iopub.status.idle":"2026-05-10T03:41:52.322141Z","shell.execute_reply.started":"2026-05-10T03:41:52.316231Z","shell.execute_reply":"2026-05-10T03:41:52.321002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# gaussian distribution for patch overlap\ndef get_gaussian_weight(kernel_size, sigma=1, muu=0):\n    x, y, z = torch.meshgrid(\n        torch.linspace(-kernel_size[0]//2, kernel_size[0]//2, kernel_size[0]),\n        torch.linspace(-kernel_size[1]//2, kernel_size[1]//2, kernel_size[1]), \n        torch.linspace(-kernel_size[2]//2, kernel_size[2]//2, kernel_size[2])\n        )\n    mean = torch.Tensor([muu, muu, muu])  # Mean of the Gaussian\n    std_dev = torch.Tensor([sigma, sigma, sigma])  # Standard deviation along each axis\n\n    gaussian = torch.exp(\n        -0.5 * (\n            ((x - mean[0]) / std_dev[0])**2 +\n            ((y - mean[1]) / std_dev[1])**2 +\n            ((z - mean[2]) / std_dev[2])**2\n        )\n    )\n    \n    return gaussian\n\nsigma = PATCH_SIZE[1]//2-0\nweight = get_gaussian_weight(PATCH_SIZE, sigma, 0).to('cuda')\n\n# -- check gaussian dist visually\nimport matplotlib.pyplot as plt\nplt.imshow(weight[PATCH_SIZE[0]//2,:,:].cpu().numpy(), cmap='viridis', interpolation='nearest', vmin=0, vmax=1)\nplt.colorbar()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:52.323499Z","iopub.execute_input":"2026-05-10T03:41:52.323871Z","iopub.status.idle":"2026-05-10T03:41:52.877226Z","shell.execute_reply.started":"2026-05-10T03:41:52.323821Z","shell.execute_reply":"2026-05-10T03:41:52.876427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_patch_starts(\n    dimension_size: int, \n    patch_size: int, \n    overlap_factor=2\n    ) -> List[int]:\n    \"\"\"\n    Calculate the starting positions of patches along a single dimension\n    with minimal overlap to cover the entire dimension.\n    \n    Parameters:\n    -----------\n    dimension_size : int\n        Size of the dimension\n    patch_size : int\n        Size of the patch in this dimension\n        \n    Returns:\n    --------\n    List[int]\n        List of starting positions for patches\n    \"\"\"\n    if dimension_size <= patch_size:\n        return [0]\n        \n    # Calculate number of patches needed\n    n_patches = np.ceil(dimension_size / patch_size) * overlap_factor\n    \n    if n_patches == 1:\n        return [0]\n    \n    # Calculate overlap\n    total_overlap = (n_patches * patch_size - dimension_size) / (n_patches - 1)\n    \n    # Generate starting positions\n    positions = []\n    for i in range(int(n_patches)):\n        pos = int(i * (patch_size - total_overlap))\n        if pos + patch_size > dimension_size:\n            pos = dimension_size - patch_size\n        if pos not in positions:  # Avoid duplicates\n            positions.append(pos)\n    \n    return positions\n\ndef extract_3d_patches_minimal_overlap(arrays: List[np.ndarray], patch_size: list) -> Tuple[List[np.ndarray], List[Tuple[int, int, int]]]:\n    \"\"\"\n    Extract 3D patches from multiple arrays with minimal overlap to cover the entire array.\n    \n    Parameters:\n    -----------\n    arrays : List[np.ndarray]\n        List of input arrays, each with shape (m, n, l)\n    patch_size : list\n        Size of cubic patches (a x a x a)\n        \n    Returns:\n    --------\n    patches : List[np.ndarray]\n        List of all patches from all input arrays\n    coordinates : List[Tuple[int, int, int]]\n        List of starting coordinates (x, y, z) for each patch\n    \"\"\"\n    if not arrays or not isinstance(arrays, list):\n        raise ValueError(\"Input must be a non-empty list of arrays\")\n    \n    # Verify all arrays have the same shape\n    shape = arrays[0].shape\n    if not all(arr.shape == shape for arr in arrays):\n        raise ValueError(\"All input arrays must have the same shape\")\n    \n    if min(patch_size) > min(shape):\n        raise ValueError(f\"patch_size ({min(patch_size)}) must be smaller than smallest dimension {min(shape)}\")\n    \n    m, n, l = shape\n    patches = []\n    coordinates = []\n    \n    # Calculate starting positions for each dimension\n    x_starts = calculate_patch_starts(m, patch_size[0], OVERLAP[0])\n    y_starts = calculate_patch_starts(n, patch_size[1], OVERLAP[1])\n    z_starts = calculate_patch_starts(l, patch_size[2], OVERLAP[2])\n    \n    # Extract patches from each array\n    for arr in arrays:\n        for x in x_starts:\n            for y in y_starts:\n                for z in z_starts:\n                    patch = arr[\n                        x:x + patch_size[0],\n                        y:y + patch_size[1],\n                        z:z + patch_size[2]\n                    ]\n                    patches.append(patch)\n                    coordinates.append((x, y, z))\n    \n    return patches, coordinates\n\ndef dict_to_df(coord_dict, experiment_name):\n    \"\"\"\n    Convert dictionary of coordinates to pandas DataFrame.\n    \n    Parameters:\n    -----------\n    coord_dict : dict\n        Dictionary where keys are labels and values are Nx3 coordinate arrays\n        \n    Returns:\n    --------\n    pd.DataFrame\n        DataFrame with columns ['x', 'y', 'z', 'label']\n    \"\"\"\n    # Create lists to store data\n    all_coords = []\n    all_labels = []\n    \n    # Process each label and its coordinates\n    for label, coords in coord_dict.items():\n        all_coords.append(coords)\n        all_labels.extend([label] * len(coords))\n    \n    # Concatenate all coordinates\n    all_coords = np.vstack(all_coords)\n    \n    df = pd.DataFrame({\n        'experiment': experiment_name,\n        'particle_type': all_labels,\n        'x': all_coords[:, 0],\n        'y': all_coords[:, 1],\n        'z': all_coords[:, 2]\n    })\n\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:52.880266Z","iopub.execute_input":"2026-05-10T03:41:52.880622Z","iopub.status.idle":"2026-05-10T03:41:52.892138Z","shell.execute_reply.started":"2026-05-10T03:41:52.880593Z","shell.execute_reply":"2026-05-10T03:41:52.891286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def do_one_eval(truth, predict, threshold):\n    P=len(predict)\n    T=len(truth)\n\n    if P==0:\n        hit=[[],[]]\n        miss=np.arange(T).tolist()\n        fp=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    if T==0:\n        hit=[[],[]]\n        fp=np.arange(P).tolist()\n        miss=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    #---\n    distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n    distance = distance**2\n    distance = distance.sum(axis=2)\n    distance = np.sqrt(distance)\n    p_index, t_index = linear_sum_assignment(distance)\n\n    valid = distance[p_index, t_index] <= threshold\n    p_index = p_index[valid]\n    t_index = t_index[valid]\n    hit = [p_index.tolist(), t_index.tolist()]\n    miss = np.arange(T)\n    miss = miss[~np.isin(miss,t_index)].tolist()\n    fp = np.arange(P)\n    fp = fp[~np.isin(fp,p_index)].tolist()\n\n    metric = [P,T,len(hit[0]),len(miss),len(fp)] #for lb metric F-beta copmutation\n    return hit, fp, miss, metric\n\n\ndef compute_lb(submit_df, overlay_dir):\n\n    valid_id = list(submit_df['experiment'].unique())\n    print(valid_id)\n\n    eval_df = []\n    for id in valid_id:\n        truth = read_one_truth(id, overlay_dir) #=f'{valid_dir}/overlay/ExperimentRuns')\n        id_df = submit_df[submit_df['experiment'] == id]\n        \n        for p in PARTICLE:\n            p = dotdict(p)            \n            xyz_truth = truth[p.name]\n            xyz_predict = id_df[id_df['particle_type'] == p.name][['x', 'y', 'z']].values\n            hit, fp, miss, metric = do_one_eval(xyz_truth, xyz_predict, p.radius* 0.5)\n            eval_df.append(dotdict(\n                id=id, particle_type=p.name,\n                P=metric[0], T=metric[1], hit=metric[2], miss=metric[3], fp=metric[4],\n            )) \n    print('')\n\n    eval_df = pd.DataFrame(eval_df)\n    gb = eval_df.groupby('particle_type').agg('sum').drop(columns=['id'])\n    gb.loc[:, 'precision'] = gb['hit'] / gb['P']\n    gb.loc[:, 'precision'] = gb['precision'].fillna(0)\n    gb.loc[:, 'recall'] = gb['hit'] / gb['T']\n    gb.loc[:, 'recall'] = gb['recall'].fillna(0)\n    gb.loc[:, 'f-beta4'] = 17 * gb['precision'] * gb['recall'] / (16 * gb['precision'] + gb['recall'])\n    gb.loc[:, 'f-beta4'] = gb['f-beta4'].fillna(0)\n\n    gb = gb.sort_values('particle_type').reset_index(drop=False)\n    # https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895\n    gb.loc[:, 'weight'] = [1, 0, 2, 1, 2, 1]\n    lb_score = (gb['f-beta4'] * gb['weight']).sum() / gb['weight'].sum()\n    return gb, lb_score","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-05-10T03:41:52.893452Z","iopub.execute_input":"2026-05-10T03:41:52.893817Z","iopub.status.idle":"2026-05-10T03:41:52.908064Z","shell.execute_reply.started":"2026-05-10T03:41:52.893778Z","shell.execute_reply":"2026-05-10T03:41:52.907129Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Copick root / Test transform / Model","metadata":{}},{"cell_type":"code","source":"root = copick.from_file(COPICK_CONFIG_PATH)\n\ntest_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n])\n\nmodel = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=NUM_CLASSES,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=1,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:52.909230Z","iopub.execute_input":"2026-05-10T03:41:52.909741Z","iopub.status.idle":"2026-05-10T03:41:52.945362Z","shell.execute_reply.started":"2026-05-10T03:41:52.909716Z","shell.execute_reply":"2026-05-10T03:41:52.944679Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submit","metadata":{}},{"cell_type":"code","source":"path_to_model = sorted(glob(f'{MODEL_FOLDER}/*.pth'))[0]\nmodel.load_state_dict(torch.load(path_to_model))\nmodel.to(\"cuda\")\nmodel.eval()\n\nprint(f'[ submission with fold = {FOLD} ]')\nprint(\n    f'\\tcertainty threshold: {CERT_THRESH}\\n'\n    f'\\tblob threshold:      {BLOB_THRESH}\\n'\n    f'\\toverlap:             {OVERLAP}\\n'\n    f'\\ttta:                 {DO_TTA}\\n'\n    f'\\ttta k rotate:        {TTA_K_ROTATE}\\n'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:41:52.946613Z","iopub.execute_input":"2026-05-10T03:41:52.947341Z","iopub.status.idle":"2026-05-10T03:41:53.071392Z","shell.execute_reply.started":"2026-05-10T03:41:52.947298Z","shell.execute_reply":"2026-05-10T03:41:53.070504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    location_df = []\n    inference_time = []\n    for run in root.runs:\n\n        print(f'testing {run.name} ...')\n        start = time.time()\n\n        tomo = run.get_voxel_spacing(10)\n        tomo = tomo.get_tomogram('denoised').numpy()\n        original_shape = tomo.shape\n        tomo_patches, coordinates  = extract_3d_patches_minimal_overlap([tomo], PATCH_SIZE)\n        tomo_patched_data = [{\"image\": img} for img in tomo_patches]\n        tomo_ds = CacheDataset(data=tomo_patched_data, transform=test_transforms, cache_rate=1.0)\n\n        reconstructed = torch.zeros(\n            [NUM_CLASSES, original_shape[0], original_shape[1], original_shape[2]]\n        ).to('cuda')  # To track overlapping regions\n        count = torch.zeros(\n            [NUM_CLASSES, original_shape[0], original_shape[1], original_shape[2]]\n        ).to('cuda')\n        \n        for i in range(len(tomo_ds)):\n            \n            input_tensor = tomo_ds[i]['image'].unsqueeze(0).to(\"cuda\")\n            \n            if DO_TTA:\n                # w/o rotate\n                input_tensor = tomo_ds[i]['image'].unsqueeze(0).to('cuda')\n                model_output_tmp = model(input_tensor)\n                model_output_tmp = model_output_tmp.squeeze(0)\n                model_outputs_tta = [model_output_tmp]\n                # tta with rotate90(k=1~3)\n                for k in range(1, TTA_K_ROTATE+1):\n                    input_tensor   = tomo_ds[i]['image']\n                    rotate         = Rotate90(k=k, spatial_axes=(0, 2))\n                    rotate_inverse = Rotate90(k=4-k, spatial_axes=(0, 2))\n                    input_tensor = rotate(input_tensor)\n                    input_tensor = input_tensor.unsqueeze(0).to(\"cuda\")\n                    model_output_tmp = model(input_tensor)\n                    model_output_tmp = model_output_tmp.squeeze(0)\n                    model_output_tmp = rotate_inverse(model_output_tmp)\n                    model_outputs_tta.append(model_output_tmp)\n                model_output = torch.stack(model_outputs_tta, 0).mean(0)\n                model_output = model_output.unsqueeze(0)\n            else: \n                model_output = model(input_tensor)\n                \n            prob = torch.softmax(model_output[0], dim=0) #prob.shape: (7,96,96,96)\n            \n            reconstructed[\n                :, \n                coordinates[i][0]:coordinates[i][0] + PATCH_SIZE[0],\n                coordinates[i][1]:coordinates[i][1] + PATCH_SIZE[1],\n                coordinates[i][2]:coordinates[i][2] + PATCH_SIZE[2]\n            ] += prob\n\n            count[\n                :, \n                coordinates[i][0]:coordinates[i][0] + PATCH_SIZE[0],\n                coordinates[i][1]:coordinates[i][1] + PATCH_SIZE[1],\n                coordinates[i][2]:coordinates[i][2] + PATCH_SIZE[2]\n            ] += weight\n\n        reconstructed /= count\n            \n        max_probs, max_classes = torch.max(reconstructed, dim=0)\n        thresh_prob = torch.zeros_like(reconstructed)\n        thresh_max_classes = torch.zeros_like(reconstructed[0])\n        for ch in range(NUM_CLASSES):\n            max_channel_is_one = torch.where(max_classes==ch, 1, 0)\n            thresh_prob[ch] = max_probs * max_channel_is_one > CERT_THRESH[ch]\n            thresh_prob[ch] = torch.where(thresh_prob[ch]==1, ch, 0)\n            thresh_max_classes += thresh_prob[ch]\n\n        thresh_max_classes = thresh_max_classes.cpu().numpy()\n        \n        location = {}\n        for c in CLASS_IDS:\n            cc = cc3d.connected_components(thresh_max_classes == c)\n            stats = cc3d.statistics(cc)\n            zyx=stats['centroids'][1:]*10.012444 #https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895#3040071\n            zyx_large = zyx[stats['voxel_counts'][1:] > BLOB_THRESH]\n            xyz =np.ascontiguousarray(zyx_large[:,::-1])\n            location[ID_TO_NAME[c]] = xyz\n        df = dict_to_df(location, run.name)\n        location_df.append(df)\n\n        inference_time.append(time.time()-start)\n    \n    location_df = pd.concat(location_df)\n\nprint(location_df)\nlocation_df.insert(loc=0, column='id', value=np.arange(len(location_df)))\nlocation_df.to_csv(\"submission.csv\", index=False)\n\nmean_inference_time = sum(inference_time)/len(inference_time)\n\nprint('\\ninference for submission done!!')\nprint(\n    f'\\nmean inference time: {mean_inference_time:.2f} sec'\n    f'\\nestimated scoring time: {mean_inference_time*500/60/60:.2f} hr'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T03:43:50.413407Z","iopub.execute_input":"2026-05-10T03:43:50.414498Z","iopub.status.idle":"2026-05-10T03:47:32.663184Z","shell.execute_reply.started":"2026-05-10T03:43:50.414460Z","shell.execute_reply":"2026-05-10T03:47:32.662015Z"}},"outputs":[],"execution_count":null}]}