{"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":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":10347772,"sourceType":"datasetVersion","datasetId":6407621},{"sourceId":10476214,"sourceType":"datasetVersion","datasetId":6486959},{"sourceId":10624471,"sourceType":"datasetVersion","datasetId":6578264},{"sourceId":220221133,"sourceType":"kernelVersion"},{"sourceId":220714418,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Install 3rdparty\ninference by https://www.kaggle.com/code/fnands/baseline-unet-train-submit/notebook","metadata":{"execution":{"iopub.status.busy":"2024-12-09T04:24:59.958511Z","iopub.execute_input":"2024-12-09T04:24:59.958788Z","iopub.status.idle":"2024-12-09T04:24:59.969213Z","shell.execute_reply.started":"2024-12-09T04:24:59.95876Z","shell.execute_reply":"2024-12-09T04:24:59.967921Z"},"_kg_hide-input":true}},{"cell_type":"code","source":"root_dir = \"/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 asciitree-0.3.3-py3-none-any.whl\n# !pip install -q --no-index --find-links {root_dir} --requirement {root_dir}/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:42:36.345348Z","iopub.execute_input":"2025-02-04T07:42:36.345753Z","iopub.status.idle":"2025-02-04T07:43:54.993521Z","shell.execute_reply.started":"2025-02-04T07:42:36.345720Z","shell.execute_reply":"2025-02-04T07:43:54.992700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q --no-index --find-links {root_dir} --requirement {root_dir}/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:43:54.995302Z","iopub.execute_input":"2025-02-04T07:43:54.995587Z","iopub.status.idle":"2025-02-04T07:44:17.512530Z","shell.execute_reply.started":"2025-02-04T07:43:54.995559Z","shell.execute_reply":"2025-02-04T07:44:17.511736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/tensorrt-10-1-0/nvidia_cuda_runtime_cu12-12.2.140-py3-none-manylinux1_x86_64.whl\n!pip install /kaggle/input/tensorrt-10-1-0/tensorrt_cu12_bindings-10.1.0-cp310-none-manylinux_2_17_x86_64.whl\n!pip install /kaggle/input/tensorrt-10-1-0/tensorrt_cu12_libs-10.1.0-py2.py3-none-manylinux_2_17_x86_64.whl\n!pip install /kaggle/input/tensorrt-10-1-0/tensorrt_cu12-10.1.0-py2.py3-none-any.whl\n!pip install /kaggle/input/tensorrt-10-1-0/tensorrt-10.1.0-py2.py3-none-any.whl\n!pip install /kaggle/input/tensorrt-10-1-0/polygraphy-0.49.14-py2.py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:44:17.513851Z","iopub.execute_input":"2025-02-04T07:44:17.514131Z","iopub.status.idle":"2025-02-04T07:48:43.291470Z","shell.execute_reply.started":"2025-02-04T07:44:17.514104Z","shell.execute_reply":"2025-02-04T07:48:43.290465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cp -r /kaggle/input/tensorrt-10-1-0/torch2trt-master /kaggle/working/torch2trt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:48:43.292981Z","iopub.execute_input":"2025-02-04T07:48:43.293330Z","iopub.status.idle":"2025-02-04T07:48:45.328350Z","shell.execute_reply.started":"2025-02-04T07:48:43.293300Z","shell.execute_reply":"2025-02-04T07:48:45.327226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/working/torch2trt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:48:45.331699Z","iopub.execute_input":"2025-02-04T07:48:45.332362Z","iopub.status.idle":"2025-02-04T07:49:37.489918Z","shell.execute_reply.started":"2025-02-04T07:48:45.332319Z","shell.execute_reply":"2025-02-04T07:49:37.488830Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Import 3rdparty","metadata":{}},{"cell_type":"code","source":"import os \nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nfrom tqdm import tqdm\n\nimport torch\nimport copick\n\nfrom monai.networks.nets import UNet\nfrom monai.inferers import sliding_window_inference\nfrom monai.data import DataLoader, CacheDataset\nfrom monai.transforms import (\n    Compose, \n    NormalizeIntensityd,\n    EnsureChannelFirstd, \n    Activationsd,\n    AsDiscreted,\n    Orientationd\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:49:37.491271Z","iopub.execute_input":"2025-02-04T07:49:37.491562Z","iopub.status.idle":"2025-02-04T07:50:25.998156Z","shell.execute_reply.started":"2025-02-04T07:49:37.491535Z","shell.execute_reply":"2025-02-04T07:50:25.997469Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Set Config","metadata":{}},{"cell_type":"code","source":"cfg_list = []\nclass Config:\n    data_dir = '/kaggle/input/czii-cryo-et-object-identification'\n    train_dir = data_dir + 'train/'\n    test_dir = data_dir + 'test/'\n    \n    ### cnn backbone\n    model_name = 'Unet'\n    # spatial_size = [32,320,320]\n    # spatial_size = [40,224,224]\n    \n    # spatial_size = [48,256,256]\n    spatial_size = [180] * 3\n    # spatial_size = [56,224,224]\n    # spatial_size = [64,256,256]\n    # spatial_size = [64,192,192]\n    # spatial_size = [96,128,128]\n\n    \n    spatial_dims = 3\n    in_channels = 1\n    # channels = (48, 64, 80, 80) # small\n    channels = (48, 64, 80, 96) # middle\n    # channels = (64,80,96,112) # big\n    strides = (2, 2, 1)\n    num_res_units = 1\n    num_samples = 16\n    batch_size = 1\n    num_classes = 7\n    dropout = 0.1\n    \nCFG = Config\ncfg_list.append(CFG)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:25.999143Z","iopub.execute_input":"2025-02-04T07:50:25.999960Z","iopub.status.idle":"2025-02-04T07:50:26.005780Z","shell.execute_reply.started":"2025-02-04T07:50:25.999929Z","shell.execute_reply":"2025-02-04T07:50:26.004489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nMODE = \"submit\"\nprint(f\"MODE: {MODE}\")\n\nif MODE=='local':\n    valid_dir =f'{CFG.data_dir}/train'\n    # valid_id = ['TS_5_4','TS_6_4' ] #f0\n    valid_id = ['TS_69_2', ] #f1\n    # valid_id = ['TS_6_6', 'TS_73_6'] #f2\n    # valid_id = ['TS_86_3', 'TS_99_9' ] #f3\n\n    \nif MODE=='submit':\n    valid_dir =f'{CFG.data_dir}/test' \n    valid_id = glob.glob(f'{valid_dir}/static/ExperimentRuns/*')\n    valid_id = [f.split('/')[-1] for f in valid_id]\n    \nprint('valid_id:',len(valid_id), valid_id)\nprint('MODE:', MODE)\nprint('SETTING OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:26.007059Z","iopub.execute_input":"2025-02-04T07:50:26.007416Z","iopub.status.idle":"2025-02-04T07:50:26.049376Z","shell.execute_reply.started":"2025-02-04T07:50:26.007370Z","shell.execute_reply":"2025-02-04T07:50:26.048643Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PL Model","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\nclass CustomPLModel(pl.LightningModule):\n    def __init__(self,cfg):\n        super(CustomPLModel,self).__init__()\n        self.model = UNet(\n            spatial_dims=cfg.spatial_dims,\n            in_channels=cfg.in_channels,\n            out_channels=cfg.num_classes,\n            channels=cfg.channels,\n            strides=cfg.strides,\n            num_res_units=cfg.num_res_units,\n        )\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:26.050308Z","iopub.execute_input":"2025-02-04T07:50:26.050542Z","iopub.status.idle":"2025-02-04T07:50:27.207945Z","shell.execute_reply.started":"2025-02-04T07:50:26.050513Z","shell.execute_reply":"2025-02-04T07:50:27.207028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load model weights","metadata":{}},{"cell_type":"code","source":"from torch2trt import TRTModule\nimport tensorrt as trt\nfrom cuda import cudart\n\nLOGGER = trt.Logger(trt.Logger.INFO)\ntrt.init_libnvinfer_plugins(LOGGER, \"\")\n\nclass Net:\n    def __init__(self, weights, device=0):\n        self.device = device\n        \n        # 推荐使用 torch.cuda.set_device(self.device)\n        # 或者跟你原来的一样 cudart.cudaSetDevice(device)\n        torch.cuda.set_device(self.device)\n        \n        self.runtime = trt.Runtime(LOGGER)\n        with open(weights, \"rb\") as f:\n            self.engine = self.runtime.deserialize_cuda_engine(f.read())\n        \n        # 创建 TRTModule\n        self.trt_model = TRTModule(\n            input_names=['images'],\n            output_names=['output'],\n            engine=self.engine\n        )\n    \n    # 如果没有必要，尽量不要加 __del__\n    # def __del__(self):\n    #     del self.trt_model\n\n    def __call__(self, img):\n        # 确保输入也在同一块 GPU\n        if img.device.index != self.device:\n            img = img.cuda(self.device)\n        \n        with torch.no_grad():\n            output = self.trt_model(img)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.209278Z","iopub.execute_input":"2025-02-04T07:50:27.209952Z","iopub.status.idle":"2025-02-04T07:50:27.313462Z","shell.execute_reply.started":"2025-02-04T07:50:27.209910Z","shell.execute_reply":"2025-02-04T07:50:27.312646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nclass TRTEnsembleModel(nn.Module):\n    def __init__(self, models, weights, size):\n        super(TRTEnsembleModel, self).__init__()\n        self.models = []\n        self.weights = weights\n        # for model_config in models:\n        #     self.models.append(Net(weights=model_config['engine_path'], device=model_config['device']))\n        self.models = models\n        self.device = 'cuda:0' if models[0]['device'] == 0 else 'cuda:1'\n        self.size = size\n    def predict_tomo(self,model,size,data):\n        model = Net(weights=model['engine_path'], device=model['device'])\n        pred = 0\n        tta_ct = 0\n        sw_batch_size = 2\n        device = self.device\n        TTA = True\n        # TTA_LIST = [[3],[4],[3,4]]\n        TTA_LIST = [[2],[3],[4],[3,4]]\n        weights = [1.0]\n        ROTATE_TTA = True\n        ROTATE_K_LIST = [1,2,3]\n        OVERLAP_T = 0.15\n        out_device = device\n        with torch.amp.autocast(device, enabled=True), torch.cuda.device(device):\n            with torch.no_grad():\n                tomogram = data['image'].to(device) \n                pred += sliding_window_inference(\n                        inputs=tomogram,\n                        roi_size=size,\n                        sw_batch_size=sw_batch_size, \n                        sw_device=device,\n                        predictor=model,\n                        overlap=OVERLAP_T,\n                        device=out_device,\n                        progress=False\n                    )\n                tta_ct += 1\n                torch.cuda.empty_cache()\n                if TTA:\n                    for dims in TTA_LIST:\n                        t_pred = sliding_window_inference(\n                                                inputs=torch.flip(tomogram, dims=dims),\n                                                roi_size=size,\n                                                sw_batch_size=sw_batch_size,\n                                                sw_device=device,                \n                                                predictor=model,\n                                                overlap=OVERLAP_T,\n                                                device=out_device,\n                                                progress=False\n                                            )\n    \n                        pred += torch.flip(t_pred, dims=dims)\n                        tta_ct += 1\n                        torch.cuda.empty_cache()\n                if ROTATE_TTA:\n                    for k in ROTATE_K_LIST:\n                        t_pred = sliding_window_inference(\n                                        inputs=torch.rot90(tomogram, k=k,dims=[3,4]),\n                                        roi_size=size,\n                                        sw_batch_size=sw_batch_size, \n                                        sw_device=device,\n                                        predictor=model,\n                                        overlap=OVERLAP_T,\n                                        device=out_device,\n                                        progress=False\n                                    )\n                        pred += torch.rot90(t_pred, k=-k,dims=[3,4])\n                        tta_ct += 1\n                        torch.cuda.empty_cache()\n        pred /= tta_ct \n        del model\n        torch.cuda.empty_cache()\n        return pred\n    \n    def forward(self, x):\n        with torch.no_grad():\n            outputs = [self.predict_tomo(model,size,x) * weight for model, weight, size in zip(self.models, self.weights, self.size)]\n        return sum(outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.314788Z","iopub.execute_input":"2025-02-04T07:50:27.315049Z","iopub.status.idle":"2025-02-04T07:50:27.325619Z","shell.execute_reply.started":"2025-02-04T07:50:27.315023Z","shell.execute_reply":"2025-02-04T07:50:27.324837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weight_0 = [\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_0.engine', 'device': 0},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_1.engine', 'device': 0},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_2.engine', 'device': 0},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_3.engine', 'device': 0},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_4.engine', 'device': 0},\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.326470Z","iopub.execute_input":"2025-02-04T07:50:27.326744Z","iopub.status.idle":"2025-02-04T07:50:27.343055Z","shell.execute_reply.started":"2025-02-04T07:50:27.326704Z","shell.execute_reply":"2025-02-04T07:50:27.342428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weight_1 = [\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_0.engine', 'device': 1},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_1.engine', 'device': 1},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_2.engine', 'device': 1},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_3.engine', 'device': 1},\n    {'engine_path': '/kaggle/input/make-onnx-high-lb-model/trt/model_4.engine', 'device': 1},\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.344040Z","iopub.execute_input":"2025-02-04T07:50:27.344315Z","iopub.status.idle":"2025-02-04T07:50:27.359641Z","shell.execute_reply.started":"2025-02-04T07:50:27.344287Z","shell.execute_reply":"2025-02-04T07:50:27.358823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights = [0.282, 0.128, 0.282, 0.282, 0.026]\nsize = [\n    [180]*3,\n    [180]*3,\n    [180]*3,\n    [176]*3,\n    [176]*3\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.362480Z","iopub.execute_input":"2025-02-04T07:50:27.362763Z","iopub.status.idle":"2025-02-04T07:50:27.370492Z","shell.execute_reply.started":"2025-02-04T07:50:27.362738Z","shell.execute_reply":"2025-02-04T07:50:27.369919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trt_model_0 = TRTEnsembleModel(weight_0,weights,size)\ntrt_model_1 = TRTEnsembleModel(weight_1,weights,size)\n# trt_model_2 = Net(weights='/kaggle/input/make-trt-main-ensemble-original/ensemble_model2/load_3_models_ensemble.engine', device=0)\n# trt_model_3 = Net(weights='/kaggle/input/make-trt-main-ensemble-original/ensemble_model2/load_3_models_ensemble.engine', device=1)\nmodels_device0_0 = [trt_model_0]\nmodels_device1_1 = [trt_model_1]\n# models_device0_2 = [trt_model_2]\n# models_device1_3 = [trt_model_3]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.371259Z","iopub.execute_input":"2025-02-04T07:50:27.371464Z","iopub.status.idle":"2025-02-04T07:50:27.380733Z","shell.execute_reply.started":"2025-02-04T07:50:27.371443Z","shell.execute_reply":"2025-02-04T07:50:27.379949Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Helpers","metadata":{}},{"cell_type":"code","source":"def 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    \n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.381871Z","iopub.execute_input":"2025-02-04T07:50:27.382438Z","iopub.status.idle":"2025-02-04T07:50:27.392993Z","shell.execute_reply.started":"2025-02-04T07:50:27.382388Z","shell.execute_reply":"2025-02-04T07:50:27.392403Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load test data","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport json\nimport zarr\nfrom timeit import default_timer as timer\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t)/60\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min)\n\n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\n\telse:\n\t\traise NotImplementedError\n\nPARTICLE= [\n    {\n        \"name\": \"apo-ferritin\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"4V1W\",\n        \"label\": 1,\n        \"color\": [0, 255, 0, 0],\n        \"radius\": 60,\n        \"map_threshold\": 0.0418\n    },\n    {\n        \"name\": \"beta-amylase\",\n        \"difficulty\": 'ignore',\n        \"pdb_id\": \"1FA2\",\n        \"label\": 2,\n        \"color\": [0, 0, 255, 255],\n        \"radius\": 65,\n        \"map_threshold\": 0.035\n    },\n    {\n        \"name\": \"beta-galactosidase\",\n        \"difficulty\": 'hard',\n        \"pdb_id\": \"6X1Q\",\n        \"label\": 3,\n        \"color\": [0, 255, 0, 255],\n        \"radius\": 90,\n        \"map_threshold\": 0.0578\n    },\n    {\n        \"name\": \"ribosome\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"6EK0\",\n        \"label\": 4,\n        \"color\": [0, 0, 255, 0],\n        \"radius\": 150,\n        \"map_threshold\": 0.0374\n    },\n    {\n        \"name\": \"thyroglobulin\",\n        \"difficulty\": 'hard',\n        \"pdb_id\": \"6SCJ\",\n        \"label\": 5,\n        \"color\": [0, 255, 255, 0],\n        \"radius\": 130,\n        \"map_threshold\": 0.0278\n    },\n    {\n        \"name\": \"virus-like-particle\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"6N4V\",\n        \"label\": 6,\n        \"color\": [0, 0, 0, 255],\n        \"radius\": 135,\n        \"map_threshold\": 0.201\n    }\n]\n\nPARTICLE_COLOR=[[0,0,0]]+[\n    PARTICLE[i]['color'][1:] for i in range(6)\n]\nPARTICLE_NAME=['none']+[\n    PARTICLE[i]['name'] for i in range(6)\n]\n\n'''\n(184, 630, 630)  \n(92, 315, 315)  \n(46, 158, 158)  \n'''\n\ndef read_one_data(id, static_dir,min=1,max=99,precision=\"16bit\"):\n    zarr_dir = f'{static_dir}/{id}/VoxelSpacing10.000'\n    zarr_file = f'{zarr_dir}/denoised.zarr'\n    zarr_data = zarr.open(zarr_file, mode='r')\n    volume = zarr_data[0][:]\n    if precision == \"16bit\":\n        min = np.percentile(volume,min)\n        max = np.percentile(volume,max)\n        volume = np.clip(volume,min,max)\n        volume = (volume - min) / (max - min)\n        volume = (volume * 65535).astype(np.uint16)\n    elif precision == \"8bit\":\n        min = np.percentile(volume,min)\n        max = np.percentile(volume,max)\n        volume = np.clip(volume,min,max)\n        volume = (volume - min) / (max - min)\n        volume = (volume * 255).astype(np.uint8)\n    else:\n        max = volume.max()\n        min = volume.min()\n        volume = (volume - min) / (max - min)\n        volume = volume.astype(np.float16)\n    return volume\n\n\ndef read_one_truth(id, overlay_dir):\n    location={}\n\n    json_dir = f'{overlay_dir}/{id}/Picks'\n    for p in PARTICLE_NAME[1:]:\n        json_file = f'{json_dir}/{p}.json'\n\n        with open(json_file, 'r') as f:\n            json_data = json.load(f)\n\n        num_point = len(json_data['points'])\n        loc = np.array([list(json_data['points'][i]['location'].values()) for i in range(num_point)])\n        location[p] = loc\n\n    return location\n\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)\n\ndef filter_probability(probability,threshold_list):\n    _,D,H,W = probability.shape\n\n    for p in PARTICLE:\n        p = dotdict(p)\n        l = p.label\n        probability[l] = probability[l]>threshold_list[p.name]\n    return probability","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.394246Z","iopub.execute_input":"2025-02-04T07:50:27.394459Z","iopub.status.idle":"2025-02-04T07:50:27.417087Z","shell.execute_reply.started":"2025-02-04T07:50:27.394437Z","shell.execute_reply":"2025-02-04T07:50:27.416420Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# id_to_name = {1: \"apo-ferritin\", \n#               2: \"beta-amylase\",\n#               3: \"beta-galactosidase\", \n#               4: \"ribosome\", \n#               5: \"thyroglobulin\", \n#               6: \"virus-like-particle\"}\nid_to_name = {1: \"apo-ferritin\", \n              2: \"beta-galactosidase\", \n              3: \"ribosome\", \n              4: \"thyroglobulin\", \n              5: \"virus-like-particle\"}\nPARTICLE_THRESHOLDS={ \n        'apo-ferritin': 0.8,\n        'beta-amylase': 0.8,\n        'beta-galactosidase': 0.8,\n        'ribosome': 0.8,\n        'thyroglobulin': 0.8,\n        'virus-like-particle': 0.8,\n    }\n\nCLASSES = [1, 2, 3, 4, 5]\nradius_thresh = 0.2\n# BLOB_THRESHOLDS = {\n#     'apo-ferritin': np.pi * 6 * 6 * radius_thresh,\n#     'beta-amylase': np.pi * 6.5 * 6.5 * radius_thresh,\n#     'beta-galactosidase': np.pi * 9 * 9 * 0.2,\n#     'ribosome': np.pi * 15 * 15 * radius_thresh,\n#     'thyroglobulin': np.pi * 13 * 13 * 0.2,\n#     'virus-like-particle': np.pi * 13.5 * 13.5 * radius_thresh,\n# }\nBLOB_THRESHOLDS = {1: 2, 2: 33, 3: 78, 4: 42, 5: 400}\nTTA = True\n# TTA_LIST = [[3],[4],[3,4]]\nTTA_LIST = [[2],[3],[4],[3,4]]\nweights = [1.0]\nROTATE_TTA = True\nROTATE_K_LIST = [1,2,3]\nOVERLAP_T = 0.15\n\n\nPRECISION=\"\"\nprint(f\"PARTICLE_THRESHOLDS: {PARTICLE_THRESHOLDS}, \\\nBLOB_THRESHOLDS: {BLOB_THRESHOLDS}, TTA：{TTA}, TTA_LIST: {TTA_LIST},\\\nROTATE_TTA: {ROTATE_TTA}, ROTATE_K_LIST: {ROTATE_K_LIST},\\\nPRECISION: {PRECISION}, radius_thresh: {radius_thresh},weights: {weights}.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.418081Z","iopub.execute_input":"2025-02-04T07:50:27.418778Z","iopub.status.idle":"2025-02-04T07:50:27.437146Z","shell.execute_reply.started":"2025-02-04T07:50:27.418740Z","shell.execute_reply":"2025-02-04T07:50:27.436353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Non-random transforms to be cached\ninference_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n#     Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.438139Z","iopub.execute_input":"2025-02-04T07:50:27.438402Z","iopub.status.idle":"2025-02-04T07:50:27.462222Z","shell.execute_reply.started":"2025-02-04T07:50:27.438376Z","shell.execute_reply":"2025-02-04T07:50:27.461604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_tomo(cfg,model,test_loader,device):\n    pred = 0\n    tta_ct = 0\n    sw_batch_size = 2\n    device = 'cuda:0' if device == 0 else 'cuda:1'\n    with torch.amp.autocast(device, enabled=True), torch.cuda.device(device):\n        with torch.no_grad():\n            for data in test_loader:\n                pred = model(data)\n    return pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.462981Z","iopub.execute_input":"2025-02-04T07:50:27.463205Z","iopub.status.idle":"2025-02-04T07:50:27.478704Z","shell.execute_reply.started":"2025-02-04T07:50:27.463182Z","shell.execute_reply":"2025-02-04T07:50:27.477921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\nimport cc3d\nimport torch.nn.functional as F\nimport gc\n\ndef inference(cfg_list,models_device0_0,models_device1_1,weights, valid_ids):\n    start_timer = timer()\n    \n    ids_1 = valid_ids[0::2]  \n    ids_2  = valid_ids[1::2]\n    # ids_3 = valid_ids[2::4]  \n    # ids_4  = valid_ids[3::4]  \n\n    all_location_df = []\n\n    def process_ids_on_device(ids_subset, device_idx, models):\n        sub_location_df = []\n        torch.cuda.set_device(device_idx)\n        device = 'cuda:0' if device_idx == 0 else 'cuda:1'\n        with torch.no_grad():\n            threshold = list(PARTICLE_THRESHOLDS.values())\n            threshold = torch.tensor(threshold, device=device).reshape(6, 1, 1, 1)\n        with torch.no_grad():\n            with torch.cuda.device(device),torch.amp.autocast(device_type=device):\n                for idx, image_id in enumerate(ids_subset):\n                    torch.cuda.empty_cache()\n                    # print(f\"[Thread GPU={device_idx}] Start inference tomogram index: {idx}, image_id: {image_id}\")\n\n                    # 1) 读数据\n                    tomogram = read_one_data(\n                        image_id,\n                        static_dir=f'{valid_dir}/static/ExperimentRuns',\n                        precision=PRECISION\n                    )\n\n                    test_dataset = {\"image\": tomogram}\n                    test_ds = CacheDataset(\n                        data=[test_dataset],\n                        transform=inference_transforms,\n                        runtime_cache=True,\n                        cache_rate=1.0,\n                        progress=False,\n                        num_workers=16\n                    )\n\n                    test_loader = DataLoader(\n                        test_ds,\n                        batch_size=1,\n                        shuffle=False,\n                        pin_memory=torch.cuda.is_available()\n                    )\n\n                    # 2) 多模型\n                    ensemble_pred = 0\n                    for (cfg, model, weight) in zip(cfg_list, models, weights):\n                        pred = predict_tomo(cfg, model, test_loader, device_idx)\n                        ensemble_pred += pred * weight\n                        torch.cuda.empty_cache()\n\n                    # 3) 后处理\n                    probability0 = torch.softmax(ensemble_pred[0], dim=0)\n                    # probability0 = pred_probs#[1:]\n                    binary0 = (probability0 > threshold).cpu().data.numpy()\n\n                    location = {}\n                    for c in CLASSES:\n                        cc = cc3d.connected_components(binary0[c])\n                        stats = cc3d.statistics(cc)\n                        zyx = stats['centroids'][1:] * 10\n                        zyx_large = zyx[stats['voxel_counts'][1:] > BLOB_THRESHOLDS[c]]\n                        xyz = np.ascontiguousarray(zyx_large[:, ::-1])\n                        location[id_to_name[c]] = xyz\n\n                    df = dict_to_df(location, image_id)\n                    sub_location_df.append(df)\n                    del tomogram,test_dataset,test_ds,test_loader,pred,ensemble_pred,binary0,probability0\n                    torch.cuda.empty_cache()\n                    gc.collect()\n                    # print(f\"[Thread GPU={device_idx}] Done inference for image_id: {image_id}\")\n        \n        return sub_location_df\n\n   \n    with ThreadPoolExecutor(max_workers=2) as executor:\n        future_1 = executor.submit(process_ids_on_device, ids_1, 0, models_device0_0)\n        future_2  = executor.submit(process_ids_on_device, ids_2, 1, models_device1_1)\n        # future_3 = executor.submit(process_ids_on_device, ids_3, 0, models_device0_2)\n        # future_4  = executor.submit(process_ids_on_device, ids_4, 1, models_device1_3)\n\n        location_df_1 = future_1.result()  \n        location_df_2  = future_2.result()\n        # location_df_3  = future_3.result()   \n        # location_df_4  = future_4.result()   \n\n        all_location_df = location_df_1 + location_df_2 #+ location_df_3 + location_df_4\n\n    total_time = timer() - start_timer\n    num_volume = len(valid_ids)\n    print(f\"\\nDone! Processed {num_volume} volumes in {time_to_str(total_time, 'min')}\")\n    print(f'Total time for 500 volumes:', time_to_str(total_time/(num_volume+1)*500, 'min'))\n    submit_df = pd.concat(all_location_df, ignore_index=True)\n    return submit_df\nsubmit_df = inference(cfg_list,models_device0_0,models_device1_1,weights,valid_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:50:27.479874Z","iopub.execute_input":"2025-02-04T07:50:27.480113Z","iopub.status.idle":"2025-02-04T07:54:15.006317Z","shell.execute_reply.started":"2025-02-04T07:50:27.480089Z","shell.execute_reply":"2025-02-04T07:54:15.005417Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"!rm -rf /kaggle/working/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:54:15.007969Z","iopub.execute_input":"2025-02-04T07:54:15.008814Z","iopub.status.idle":"2025-02-04T07:54:16.084168Z","shell.execute_reply.started":"2025-02-04T07:54:15.008784Z","shell.execute_reply":"2025-02-04T07:54:16.083021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df.insert(loc=0, column='id', value=np.arange(len(submit_df)))\nsubmit_df.to_csv(\"submission.csv\", index=False)\nprint(submit_df.shape)\nsubmit_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:54:16.085779Z","iopub.execute_input":"2025-02-04T07:54:16.086091Z","iopub.status.idle":"2025-02-04T07:54:16.135512Z","shell.execute_reply.started":"2025-02-04T07:54:16.086063Z","shell.execute_reply":"2025-02-04T07:54:16.134544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.optimize import linear_sum_assignment\nimport cc3d\nimport matplotlib.pyplot as plt\ndef 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    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        for p in PARTICLE:\n            p = dotdict(p)\n            print('\\r', id, p.name, end='', flush=True)\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    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\n\n\n\n#debug\nif 1:\n    if MODE=='local':\n    #if 1:\n        submit_df=pd.read_csv(\n           'submission.csv'\n            # '/kaggle/input/hengck-czii-cryo-et-weights-01/submission.csv'\n        )\n        gb, lb_score = compute_lb(submit_df, f'{valid_dir}/overlay/ExperimentRuns')\n        print(gb)\n        print(f'lb_score:{lb_score:.4f}') #0.8246\n        print('')\n\n\n        #show one ----------------------------------\n        fig = plt.figure(figsize=(18, 8))\n\n        id = valid_id[0]\n        truth = read_one_truth(id,overlay_dir=f'{valid_dir}/overlay/ExperimentRuns')\n\n        submit_df = submit_df[submit_df['experiment']==id]\n        for p in PARTICLE:\n            p = dotdict(p)\n            xyz_truth = truth[p.name]\n            xyz_predict = submit_df[submit_df['particle_type']==p.name][['x','y','z']].values\n            hit, fp, miss, _ = do_one_eval(xyz_truth, xyz_predict, p.radius)\n            print(id, p.name)\n            print('\\t num truth   :',len(xyz_truth) )\n            print('\\t num predict :',len(xyz_predict) )\n            print('\\t num hit  :',len(hit[0]) )\n            print('\\t num fp   :',len(fp) )\n            print('\\t num miss :',len(miss) )\n\n            ax = fig.add_subplot(2, 3, p.label, projection='3d')\n            if hit[0]:\n                pt = xyz_predict[hit[0]]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.5, color='r')\n                pt = xyz_truth[hit[1]]\n                ax.scatter(pt[:,0], pt[:,1], pt[:,2], s=80, facecolors='none', edgecolors='r')\n            if fp:\n                pt = xyz_predict[fp]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=1, color='k')\n            if miss:\n                pt = xyz_truth[miss]\n                ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, alpha=1, facecolors='none', edgecolors='k')\n\n            ax.set_title(f'{p.name} ({p.difficulty})')\n\n        plt.tight_layout()\n        plt.show()\n        \n        #--- \n        zz=0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T07:54:16.136584Z","iopub.execute_input":"2025-02-04T07:54:16.136839Z","iopub.status.idle":"2025-02-04T07:54:16.156292Z","shell.execute_reply.started":"2025-02-04T07:54:16.136815Z","shell.execute_reply":"2025-02-04T07:54:16.155385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}