{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":10702405,"sourceType":"datasetVersion","datasetId":6632550},{"sourceId":206640467,"sourceType":"kernelVersion"},{"sourceId":211097053,"sourceType":"kernelVersion"},{"sourceId":219823686,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Install Runtime Dependencies","metadata":{"_uuid":"a7676fee-0f0d-4465-be9d-0228ea617636","_cell_guid":"d9af97e2-a3d7-45df-8337-197caa46ff99","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"deps_path = '/kaggle/input/czii-cryoet-dependencies'\n!cp -r '/kaggle/input/hengck-czii-cryo-et-01/wheel_file' '/kaggle/working/'\n!pip install /kaggle/working/wheel_file/asciitree-0.3.3/asciitree-0.3.3\n!pip install --no-index --find-links=/kaggle/working/wheel_file zarr\n!pip install --no-index --find-links=/kaggle/working/wheel_file connected-components-3d\n!pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt\n\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 {deps_path} --requirement {deps_path}/requirements.txt\n!cp -r /kaggle/input/czii-3d-unet-inference/cucim_download cucim_download\n!pip install --no-index --find-links=/kaggle/working/cucim_download /kaggle/working/cucim_download/cucim_cu12-24.12.0-cp310-cp310-manylinux_2_28_x86_64.whl","metadata":{"_uuid":"ef820858-4ebb-4a41-ba4e-de0f01de91a4","_cell_guid":"55fc8256-f3aa-4853-b3ea-3cdd40cf6c36","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Import Runtime Dependencies","metadata":{}},{"cell_type":"code","source":"from typing import List, Tuple, Union, Dict\nimport torch.amp as amp\nfrom monai.transforms import (\n    Compose,\n    EnsureChannelFirstd,\n    Orientationd,\n    NormalizeIntensityd,\n)\nfrom monai.inferers import SlidingWindowInferer\nfrom pytorch_lightning import LightningModule\nfrom monai.networks.nets import UNet\nimport numpy as np\nimport pandas as pd\nfrom dataclasses import dataclass\nfrom copy import deepcopy\nfrom tqdm import tqdm\nimport torch.nn as nn\nimport torch\nimport cupy as cp\nfrom cucim.skimage.feature import peak_local_max\nimport skimage.measure as measure\nfrom cucim.core.operations.morphology import distance_transform_edt\nfrom skimage.segmentation import watershed\nfrom pathlib import Path\nimport copick\nimport time","metadata":{"_uuid":"87aec7d2-39dc-4010-a1d6-accb73ed25e5","_cell_guid":"fd5499c8-5f0c-432e-be30-431de4077280","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define Copick Project","metadata":{}},{"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-galactosidase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6X1Q\",\n            \"label\": 2,\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\": 3,\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\": 4,\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\": 5,\n            \"color\": [255, 204, 153, 128],\n            \"radius\": 135,\n            \"map_threshold\": 0.201\n        },\n        {\n            \"name\": \"membrane\",\n            \"is_particle\": false,\n            \"label\": 8,\n            \"color\": [100, 100, 100, 128]\n        },\n        {\n            \"name\": \"background\",\n            \"is_particle\": false,\n            \"label\": 9,\n            \"color\": [10, 150, 200, 128]\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}\"\"\"\n\ncopick_config_path = \"/kaggle/working/test_copick.config\"\noutput_overlay = \"/kaggle/working/overlay\"\n\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utility Functions for Submissions","metadata":{}},{"cell_type":"code","source":"def dict_to_df(coord_dict: Dict[str, np.ndarray], experiment_name: str) -> pd.DataFrame:\n    \"\"\"\n    Convert dictionary of coordinates to its submission csv format\n    \"\"\"\n    all_coords = []\n    all_labels = []\n\n    for label, coords in coord_dict.items():\n        if len(coords):\n            all_coords.append(coords)\n            all_labels.extend([label] * len(coords))\n\n    if all_coords:\n        all_coords = np.vstack(all_coords)\n\n    if len(all_coords) == 0:\n        df = pd.DataFrame({\n            'experiment': experiment_name,\n            'particle_type': all_labels,\n            'x': [],\n            'y': [],\n            'z': []\n        })\n    else:\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    return df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Definitions and Post Processing Parameters","metadata":{}},{"cell_type":"code","source":"LABELS_7 = [\"apo-ferritin\", \"beta-amylase\", \"beta-galactosidase\", \"ribosome\", \"thyroglobulin\", \"virus-like-particle\"]\ninference_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n])\nPARTICLE_RADIUS_7 = {\n    \"apo-ferritin\": 60,\n    \"beta-amylase\": 65,\n    \"beta-galactosidase\": 90,\n    \"ribosome\": 150,\n    \"thyroglobulin\": 130,\n    \"virus-like-particle\": 135\n}\nblob_thresholds_7 = {\n    \"apo-ferritin\": 80.09733552923254,\n    \"beta-amylase\": 100,\n    \"beta-galactosidase\": 368.0,\n    \"ribosome\": 750.0,\n    \"thyroglobulin\": 480.0,\n    \"virus-like-particle\": 1150.3465099894624\n}\nCERTAINTY_THRESHOLDS_7 = {\n    \"apo-ferritin\": 0.1,\n    \"beta-amylase\": 0.1,\n    \"beta-galactosidase\": 0.1,\n    \"ribosome\": 0.1,\n    \"thyroglobulin\": 0.1,\n    \"virus-like-particle\": 0.1\n}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EMA(nn.Module):\n    def __init__(self, model: nn.Module, momentum=0.00001, warmup: int = None):\n        # https://www.kaggle.com/competitions/hubmap-hacking-the-human-vasculature/discussion/429060\n        # https://github.com/Lightning-AI/pytorch-lightning/issues/10914\n        super(EMA, self).__init__()\n        self.module = deepcopy(model)\n        self.module.eval()\n        self.momentum = momentum\n        self.decay = 1 - self.momentum\n        self.warmup = 0 if warmup is None else warmup\n        self.i_updates = 0\n\n    def _update(self, model, update_fn):\n        with torch.no_grad():\n            for ema_v, model_v in zip(self.module.state_dict().values(), model.state_dict().values()):\n                ema_v.copy_(update_fn(ema_v, model_v))\n\n    def update(self, model):\n        if self.i_updates < self.warmup:\n            # print(f\"setting ema <- model: {self.i_updates}/{self.warmup}\")\n            self.set(model)\n        else:\n            self._update(model, update_fn=lambda e, m: self.decay * e + (1. - self.decay) * m)\n        self.i_updates += 1\n\n    def set(self, model):\n        self._update(model, update_fn=lambda e, m: m)\n\n\nclass Model(LightningModule):\n    def __init__(\n            self,\n            spatial_dims: int = 3,\n            in_channels: int = 1,\n            out_channels: int = 7,\n            channels: Union[Tuple[int, ...], List[int]] = (32, 64, 128, 128),\n            strides: Union[Tuple[int, ...], List[int]] = (2, 2, 1),\n            num_res_units: int = 1,\n            use_ema: bool = False,\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n        self.model = UNet(\n            spatial_dims=self.hparams.spatial_dims,\n            in_channels=self.hparams.in_channels,\n            out_channels=self.hparams.out_channels,\n            channels=self.hparams.channels,\n            strides=self.hparams.strides,\n            num_res_units=self.hparams.num_res_units,\n        )\n        self.use_ema = use_ema\n        self.ema_model = EMA(self.model, 0.001, 1000)\n\n    def configure_ema_model(self):\n        # sets either self.model or self.ema_model as the one doing forward, delete the irrelevant one to save gpu space\n        if self.use_ema:\n            print(f\"model does uses ema, swapping ema_model into model\")\n            self.model = deepcopy(self.ema_model.module)\n            del self.ema_model\n        else:\n            print(f\"model does NOT uses ema, discarding ema_model\")\n            del self.ema_model\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass ModelSpec:\n    \"\"\"\n    informations required to load the model\n    \"\"\"\n    model_path: str  # checkpoint path\n    model_config: dict  # kwargs for the model initialisations\n    use_ema: bool  # whether to load the model's weight or use the model's ema weight\n\n\ndef load_models(model_specs, device_id: int):\n    \"\"\"\n    Load any number of models according to their specifications:\n    \"\"\"\n    models = []\n    for model_spec in model_specs:\n        # Create your model with the config\n        model = Model(**model_spec.model_config, use_ema=model_spec.use_ema)\n\n        # Load state dict\n        checkpoint = torch.load(model_spec.model_path)\n        if \"state_dict\" in checkpoint:\n            print(\"found state_dict key, loading state_dict\")\n            model.load_state_dict(checkpoint[\"state_dict\"], strict=False)  # strict=False for backward compat with those without ema_model\n        else:\n            print(\"can't find state_dict key, loading checkpoint directly\")\n            model.load_state_dict(checkpoint, strict=False)\n        model.configure_ema_model()\n\n        model.to(f\"cuda:{device_id}\")\n        model.eval()\n\n        models.append(model)\n    return models\n\n\nmodel_config_one = {\n    \"spatial_dims\": 3,\n    \"in_channels\": 1,\n    \"out_channels\": 7,\n    \"channels\": (32, 64, 128, 128),\n    \"strides\": (2, 2, 1),\n    \"num_res_units\": 1,\n}\nmodel_config_four = {\n    \"spatial_dims\": 3,\n    \"in_channels\": 1,\n    \"out_channels\": 7,\n    \"channels\": (32, 64, 128, 256),\n    \"strides\": (2, 2, 1),\n    \"num_res_units\": 2,\n}\nmodel_config_five = {\n    \"spatial_dims\": 3,\n    \"in_channels\": 1,\n    \"out_channels\": 7,\n    \"channels\": (32, 96, 256, 384),\n    \"strides\": (2, 2, 1),\n    \"num_res_units\": 2,\n}\n\n\nmodel_specs = [\n    ModelSpec(\"/kaggle/input/czii-final-models/large_unet-soup-folds-69-86-99.pth\", model_config_five, True),\n    ModelSpec(\"/kaggle/input/czii-final-models/medium_unet_soup.pth\", model_config_four, False),\n    ModelSpec(\"/kaggle/input/czii-final-models/tiny_unet_soup.pth\", model_config_one, False),\n    ModelSpec(\"/kaggle/input/czii-final-models/large_unet-ema-soup.pth\", model_config_five, True),\n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TTA definitions","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef tta_infer(inputs, model, inferer):\n    count = 0\n    predmask = inferer(inputs.unsqueeze(0), model)\n    count += 1\n    predmask += torch.flip(inferer(torch.flip(inputs, dims=[3]).unsqueeze(0), model), dims=[4])  # Flip prediction back\n    count += 1\n    predmask += torch.flip(inferer(torch.flip(inputs, dims=[2]).unsqueeze(0), model), dims=[3])  # Flip prediction back\n    count += 1\n    predmask += inferer(inputs.transpose(2, 3).unsqueeze(0), model).transpose(3, 4)\n    count += 1\n    return predmask / count\n\n\ndef infer_tomograph_locations_watershed_ensemble(models, tomo, exp):\n    inferer_kwargs = [\n        dict(\n            roi_size=(160, 384, 384),\n            sw_batch_size=1,\n            overlap=0.25,\n            mode=\"gaussian\",\n            padding_mode=\"reflect\",\n        ),\n    ]\n    tomo = inference_transforms({\"image\": tomo})[\"image\"].to(f\"cuda\")\n\n    # for each inferer config, run the model with tta and add logits to the ensembled score\n    predmask_accum = None\n    count = 0\n    with amp.autocast(f\"cuda\"):\n        for inferer_kwarg in inferer_kwargs:\n            for model in models:\n                print(f\"running {exp} with {inferer_kwarg}\")\n                inferer = SlidingWindowInferer(**inferer_kwarg)\n\n                this_model_pred = tta_infer(tomo, model, inferer).squeeze()\n                if predmask_accum is None:\n                    # this is the first time we have a tensor, Initialize an accumulator tensor\n                    predmask_accum = this_model_pred\n                    count = 1\n                else:\n                    # we had something initialised already, add to the accumulation\n                    predmask_accum += this_model_pred\n                    count += 1\n            torch.cuda.empty_cache()\n\n        # compute the average logits\n        predmask = predmask_accum / count\n        predmask = predmask.softmax(0)\n\n    locations = {}\n\n    # post-process via watershed segmentation and skip beta-amylase\n    predmask = cp.asarray(predmask)\n    for idx, p in tqdm(enumerate(LABELS_7)):\n        if p == \"beta-amylase\":\n            continue\n        pidx = idx + 1\n        r = PARTICLE_RADIUS_7[p] / 10\n        blob_threshold = blob_thresholds_7[p]\n        certainty_threshold = CERTAINTY_THRESHOLDS_7[p]\n\n        image = predmask[pidx].T > certainty_threshold\n\n        # 2. Compute the distance transform\n        distance = distance_transform_edt(image)\n        coords = peak_local_max(distance, min_distance=int(r), labels=image)\n        coords = cp.asnumpy(coords)\n        mask = np.zeros(distance.shape, dtype=bool)\n        mask[tuple(coords.T)] = True\n        markers = measure.label(mask)\n\n        distance = cp.asnumpy(distance)\n        image = cp.asnumpy(image)\n        labels = watershed(-distance, markers, mask=image)\n\n        regions = measure.regionprops(labels)\n        centroids = [region.centroid for region in regions if region.area >= blob_threshold]\n        locations[p] = centroids\n\n    df = dict_to_df(locations, exp)\n    df[\"x\"] *= 10.012444\n    df[\"y\"] *= 10.012444\n    df[\"z\"] *= 10.012444\n    return df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Map Directories to a Copick Project","metadata":{}},{"cell_type":"code","source":"root = copick.from_file(\"./test_copick.config\")\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Running the Models with 2 GPUs","metadata":{}},{"cell_type":"code","source":"def inference_on_runs(\n        runs: list,  # list of copick runs\n        device_number: int,\n):\n    \"\"\"\n    load models, then for each given experiments (runs), run the model on those and save out one .csv per experiment\n    \"\"\"\n    print(f\"[{device_number}]: loading models\")\n    models = load_models(model_specs, device_number)\n    print(f\"[{device_number}]: loaded models\")\n    for i_run, run in enumerate(runs):\n        start = time.time()\n\n        tomo = run.get_voxel_spacing(10)\n        tomo = tomo.get_tomogram(tomo_type).numpy()\n\n        with torch.cuda.device(f\"cuda:{device_number}\"):\n            with cp.cuda.Device(device_number):\n                loc_df = infer_tomograph_locations_watershed_ensemble(models, tomo, run.name)\n                loc_df.to_csv(f\"run_result_{device_number}_{str(i_run).zfill(5)}_{run.name}.csv\")\n\n            torch.cuda.empty_cache()\n\n        end = time.time()\n\n        print(f\"time taken: {end - start}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Split all data into 2 halves, send each half to each GPU","metadata":{}},{"cell_type":"code","source":"all_runs = root.runs\nhalf_size = int(len(all_runs) / 2)\nfirst_half = all_runs[:half_size]\nsecond_half = all_runs[half_size:]\nprint(f\"{first_half=}\")\nprint(f\"{second_half=}\")\n\nrun_in_parallel = True\n\nif run_in_parallel:\n    print(f\"running models in parallel\")\n    # our CPUs will be 100% utilised anyway, so GIL will not cause any slow downs\n    # and threading is also much easier to handle than multi processing\n    from threading import Thread\n\n    p_second_half = Thread(target=inference_on_runs, args=(second_half, 1))\n    p_first_half = Thread(target=inference_on_runs, args=(first_half, 0))\n\n    p_second_half.start()\n    p_first_half.start()\n\n    p_second_half.join()\n    p_first_half.join()\nelse:\n    print(f\"running models sequentially\")\n    inference_on_runs(second_half, 1)\n    inference_on_runs(first_half, 0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Aggregate Intermediate Results + Cleanups","metadata":{}},{"cell_type":"code","source":"csv_paths = sorted(list(Path(\".\").glob(\"run_result_*csv\")))\nprint(f\"concatting {len(csv_paths)} csvs into a sub\")\nlocation_df = []\nfor p in csv_paths:\n    df = pd.read_csv(p, index_col=0)\n    location_df.append(df)\nlocation_df = pd.concat(location_df)\nlocation_df.insert(loc=0, column='id', value=np.arange(len(location_df)))\nlocation_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -f run_result*\n!ls\n!head -n 5 submission.csv","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}