{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9840041,"sourceType":"datasetVersion","datasetId":6036295},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":11756911,"sourceType":"datasetVersion","datasetId":7380735},{"sourceId":11769950,"sourceType":"datasetVersion","datasetId":7389270},{"sourceId":11845404,"sourceType":"datasetVersion","datasetId":7442536},{"sourceId":224409375,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":3291.853074,"end_time":"2025-03-18T04:41:48.91586","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-03-18T03:46:57.062786","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dependencies","metadata":{"papermill":{"duration":0.008181,"end_time":"2025-03-18T03:46:59.816023","exception":false,"start_time":"2025-03-18T03:46:59.807842","status":"completed"},"tags":[]}},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/copick-utils copick-utils","metadata":{"papermill":{"duration":4.533548,"end_time":"2025-03-18T03:47:04.356554","exception":false,"start_time":"2025-03-18T03:46:59.823006","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:43:25.171851Z","iopub.execute_input":"2025-05-11T12:43:25.172135Z","iopub.status.idle":"2025-05-11T12:43:29.367099Z","shell.execute_reply.started":"2025-05-11T12:43:25.172103Z","shell.execute_reply":"2025-05-11T12:43:29.366208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency matplotlib tqdm copick","metadata":{"papermill":{"duration":7.009785,"end_time":"2025-03-18T03:47:11.37298","exception":false,"start_time":"2025-03-18T03:47:04.363195","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:43:29.368089Z","iopub.execute_input":"2025-05-11T12:43:29.368406Z","iopub.status.idle":"2025-05-11T12:43:36.368106Z","shell.execute_reply.started":"2025-05-11T12:43:29.368375Z","shell.execute_reply":"2025-05-11T12:43:36.367317Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency monai-weekly[mlflow]","metadata":{"papermill":{"duration":9.431744,"end_time":"2025-03-18T03:47:20.811465","exception":false,"start_time":"2025-03-18T03:47:11.379721","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:43:36.369741Z","iopub.execute_input":"2025-05-11T12:43:36.369966Z","iopub.status.idle":"2025-05-11T12:43:49.409349Z","shell.execute_reply.started":"2025-05-11T12:43:36.369945Z","shell.execute_reply":"2025-05-11T12:43:49.408331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency connected-components-3d lightning","metadata":{"papermill":{"duration":4.310163,"end_time":"2025-03-18T03:47:25.128457","exception":false,"start_time":"2025-03-18T03:47:20.818294","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:43:49.410797Z","iopub.execute_input":"2025-05-11T12:43:49.41112Z","iopub.status.idle":"2025-05-11T12:43:53.855013Z","shell.execute_reply.started":"2025-05-11T12:43:49.411092Z","shell.execute_reply":"2025-05-11T12:43:53.853782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport shutil\n\nconfig_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-train-sample-2/train/static\"\n}\"\"\"\n\ncopick_config_path = \"/kaggle/working/copick.config\"\noutput_overlay = \"/kaggle/working/overlay\"\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)\n    \n# Update the overlay\n# Define source and destination directories\nsource_dir = '/kaggle/input/czii-train-sample-2/train/overlay'\ndestination_dir = '/kaggle/working/overlay'\n\n# Walk through the source directory\nfor root, dirs, files in os.walk(source_dir):\n    # Create corresponding subdirectories in the destination\n    relative_path = os.path.relpath(root, source_dir)\n    target_dir = os.path.join(destination_dir, relative_path)\n    os.makedirs(target_dir, exist_ok=True)\n    \n    # Copy and rename each file\n    for file in files:\n        if file.startswith(\"curation_0_\"):\n            new_filename = file\n        else:\n            new_filename = f\"curation_0_{file}\"\n            \n        \n        # Define full paths for the source and destination files\n        source_file = os.path.join(root, file)\n        destination_file = os.path.join(target_dir, new_filename)\n        \n        # Copy the file with the new name\n        shutil.copy2(source_file, destination_file)\n        print(f\"Copied {source_file} to {destination_file}\")","metadata":{"papermill":{"duration":0.3179,"end_time":"2025-03-18T03:47:25.453259","exception":false,"start_time":"2025-03-18T03:47:25.135359","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:43:53.856166Z","iopub.execute_input":"2025-05-11T12:43:53.856421Z","iopub.status.idle":"2025-05-11T12:47:18.008397Z","shell.execute_reply.started":"2025-05-11T12:43:53.856399Z","shell.execute_reply":"2025-05-11T12:47:18.007766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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    RandAffined,\n    RandGaussianNoised,\n    RandStdShiftIntensityd,\n    RandShiftIntensityd,\n)\nfrom monai.networks.nets import UNet\nfrom monai.losses import DiceLoss, FocalLoss, TverskyLoss\nfrom monai.metrics import DiceMetric, ConfusionMatrixMetric\nimport mlflow\nimport mlflow.pytorch","metadata":{"papermill":{"duration":34.431565,"end_time":"2025-03-18T03:47:59.891625","exception":false,"start_time":"2025-03-18T03:47:25.46006","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:47:18.009275Z","iopub.execute_input":"2025-05-11T12:47:18.009494Z","iopub.status.idle":"2025-05-11T12:47:49.316977Z","shell.execute_reply.started":"2025-05-11T12:47:18.009477Z","shell.execute_reply":"2025-05-11T12:47:49.316273Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare the dataset","metadata":{"papermill":{"duration":0.00645,"end_time":"2025-03-18T03:47:59.905472","exception":false,"start_time":"2025-03-18T03:47:59.899022","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### 1. Get copick root","metadata":{"papermill":{"duration":0.00625,"end_time":"2025-03-18T03:47:59.918233","exception":false,"start_time":"2025-03-18T03:47:59.911983","status":"completed"},"tags":[]}},{"cell_type":"code","source":"root = copick.from_file(copick_config_path)\n\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"","metadata":{"papermill":{"duration":0.013033,"end_time":"2025-03-18T03:47:59.937709","exception":false,"start_time":"2025-03-18T03:47:59.924676","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:47:49.317756Z","iopub.execute_input":"2025-05-11T12:47:49.318741Z","iopub.status.idle":"2025-05-11T12:47:49.338309Z","shell.execute_reply.started":"2025-05-11T12:47:49.318707Z","shell.execute_reply":"2025-05-11T12:47:49.33748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2. Generate multi-class segmentation masks","metadata":{"papermill":{"duration":0.006807,"end_time":"2025-03-18T03:47:59.951218","exception":false,"start_time":"2025-03-18T03:47:59.944411","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from copick_utils.segmentation import segmentation_from_picks\nimport copick_utils.writers.write as write\nfrom collections import defaultdict\nimport numpy as np\nfrom tqdm import tqdm\n\ndef generate_segmentation_masks(root, tomo_type, copick_user_name, copick_segmentation_name):\n    target_objects = {\n        obj.name: {\n            'label': obj.label,\n            'radius': obj.radius\n        }\n        for obj in root.pickable_objects if obj.is_particle\n    }\n\n    for run in tqdm(root.runs):\n        # get tomogram\n        tomo = run.get_voxel_spacing(10)\n        tomo_array = tomo.get_tomogram(tomo_type).numpy()\n        target_mask = np.zeros(tomo_array.shape, dtype=np.uint8)\n\n        # process each particle type\n        for obj in root.pickable_objects:\n            picks = run.get_picks(object_name=obj.name, user_id=\"curation\")\n            if not picks:\n                continue\n            target_mask = segmentation_from_picks.from_picks(\n                picks[0],\n                target_mask,\n                target_objects[obj.name]['radius'] * 0.8,\n                target_objects[obj.name]['label']\n            )\n\n        # save generated segmentation\n        write.segmentation(run, target_mask, copick_user_name, name=copick_segmentation_name)\n\n# generate masks\ngenerate_segmentation_masks(root, tomo_type, copick_user_name, copick_segmentation_name)","metadata":{"papermill":{"duration":18.326325,"end_time":"2025-03-18T03:48:18.284108","exception":false,"start_time":"2025-03-18T03:47:59.957783","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:47:49.339047Z","iopub.execute_input":"2025-05-11T12:47:49.339268Z","iopub.status.idle":"2025-05-11T12:48:31.146474Z","shell.execute_reply.started":"2025-05-11T12:47:49.339249Z","shell.execute_reply":"2025-05-11T12:48:31.145609Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3. Get tomograms and their segmentaion masks","metadata":{"papermill":{"duration":0.006974,"end_time":"2025-03-18T03:48:18.29875","exception":false,"start_time":"2025-03-18T03:48:18.291776","status":"completed"},"tags":[]}},{"cell_type":"code","source":"data_dicts = []\nfor run in tqdm(root.runs):\n    tomogram = run.get_voxel_spacing(voxel_size).get_tomogram(tomo_type).numpy()\n    segmentation = run.get_segmentations(name=copick_segmentation_name, user_id=copick_user_name, voxel_size=voxel_size, is_multilabel=True)[0].numpy()\n    data_dicts.append({\"image\": tomogram, \"label\": segmentation})\n    \nprint(np.unique(data_dicts[0]['label']))","metadata":{"papermill":{"duration":6.956743,"end_time":"2025-03-18T03:48:25.262599","exception":false,"start_time":"2025-03-18T03:48:18.305856","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:48:31.14846Z","iopub.execute_input":"2025-05-11T12:48:31.148729Z","iopub.status.idle":"2025-05-11T12:49:20.262496Z","shell.execute_reply.started":"2025-05-11T12:48:31.148699Z","shell.execute_reply":"2025-05-11T12:49:20.261742Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4. Visualize the tomogram and segmentation","metadata":{"papermill":{"duration":0.007506,"end_time":"2025-03-18T03:48:25.278667","exception":false,"start_time":"2025-03-18T03:48:25.271161","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(15, 5))\n\nplt.subplot(1, 2, 1)\nplt.title('Tomogram')\nplt.imshow(data_dicts[0]['image'][100],cmap='gray')\n\nplt.subplot(1, 2, 2)\nplt.title('Painted Segmentation from Picks')\nplt.imshow(data_dicts[0]['label'][100], cmap='viridis')\n\nplt.tight_layout()\nplt.show()","metadata":{"papermill":{"duration":0.374263,"end_time":"2025-03-18T03:48:25.66075","exception":false,"start_time":"2025-03-18T03:48:25.286487","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:20.263795Z","iopub.execute_input":"2025-05-11T12:49:20.264066Z","iopub.status.idle":"2025-05-11T12:49:20.834185Z","shell.execute_reply.started":"2025-05-11T12:49:20.264043Z","shell.execute_reply":"2025-05-11T12:49:20.833204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nfrom collections import defaultdict\n\n# Root directory\noverlay_root = \"/kaggle/working/overlay/ExperimentRuns\"\n\nall_runs = sorted(os.listdir(overlay_root))\ntrain_runs = all_runs[:28]\nval_runs = all_runs[28:35]\n\ndef count_annotations(runs):\n    counts = defaultdict(int)\n    for run_name in runs:\n        run_path = os.path.join(overlay_root, run_name, \"Picks\")\n        if not os.path.exists(run_path):\n            continue\n\n        for pick_file in os.listdir(run_path):\n            if pick_file.endswith(\".json\"):\n                pick_path = os.path.join(run_path, pick_file)\n                with open(pick_path, \"r\") as f:\n                    data = json.load(f)\n                    label_name = data.get(\"pickable_object_name\", os.path.splitext(pick_file)[0])\n                    num_points = len(data.get(\"points\", []))\n                    counts[label_name] += num_points\n    return counts\n\ntrain_counts = count_annotations(train_runs)\nval_counts = count_annotations(val_runs)\n\nprint(\"=== Training Set Annotation Distribution ===\")\ntotal_train = sum(train_counts.values())\nfor label, count in sorted(train_counts.items()):\n    print(f\"{label}: {count}\")\nprint(f\"Total training annotations: {total_train}\")\n\nprint(\"\\n=== Validation Set Annotation Distribution ===\")\ntotal_val = sum(val_counts.values())\nfor label, count in sorted(val_counts.items()):\n    print(f\"{label}: {count}\")\nprint(f\"Total validation annotations: {total_val}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:20.834975Z","iopub.execute_input":"2025-05-11T12:49:20.835189Z","iopub.status.idle":"2025-05-11T12:49:20.924979Z","shell.execute_reply.started":"2025-05-11T12:49:20.835171Z","shell.execute_reply":"2025-05-11T12:49:20.924138Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5. Prepare dataloaders","metadata":{"papermill":{"duration":0.012295,"end_time":"2025-03-18T03:48:25.687605","exception":false,"start_time":"2025-03-18T03:48:25.67531","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import gc\nmy_num_samples = 16\ntrain_batch_size = 1\nval_batch_size = 1\npatch_size=[32,128,128]\n\ntrain_files, val_files = data_dicts[:28], data_dicts[28:35]\n\nprint(f\"Number of training samples: {len(train_files)}\")\nprint(f\"Number of validation samples: {len(val_files)}\")\n\n# base transforms\nbase_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\", \"label\"], axcodes=\"RAS\")\n])\n\n# transforms for training\ntrain_transforms = Compose([\n    RandCropByLabelClassesd(\n        keys=[\"image\", \"label\"],\n        label_key=\"label\",\n        spatial_size=patch_size,\n        num_classes=7,\n        num_samples=my_num_samples,\n        warn=False\n\n    ),\n    NormalizeIntensityd(keys=\"image\"),\n    RandFlipd(keys=[\"image\", \"label\"], prob=0.3, spatial_axis=0),\n    RandFlipd(keys=[\"image\", \"label\"], prob=0.3, spatial_axis=1),\n    RandFlipd(keys=[\"image\", \"label\"], prob=0.3, spatial_axis=2),\n    RandShiftIntensityd(\n        keys=[\"image\"],\n        offsets=0.10,\n        prob=0.50,\n    ),\n    RandStdShiftIntensityd(\n        keys=[\"image\"],\n        prob=0.5,\n        factors=0.1,\n    ),\n    RandGaussianNoised(keys=[\"image\"], prob=0.3, mean=0.0, std=0.1),\n    RandRotate90d(\n        keys=[\"image\", \"label\"],\n        prob=0.5,\n        max_k=3,\n        spatial_axes=[1, 2],\n    ),\n    RandAffined(\n        keys=[\"image\", \"label\"],\n        prob=0.3,\n        rotate_range=(0.1, 0.1, 0.1),\n        scale_range=(0.1, 0.1, 0.1),\n    ),\n])\n\n# apply base transforms to tarin dataset\ntrain_ds = CacheDataset(data=train_files, transform=base_transforms, cache_rate=1.0)\n\n# apply transforms\ntrain_ds = Dataset(data=train_ds, transform=train_transforms)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=train_batch_size,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available()\n)\n\n# validation transforms\nval_transforms = Compose([\n    RandCropByLabelClassesd(\n        keys=[\"image\", \"label\"],\n        label_key=\"label\",\n        spatial_size=patch_size,\n        num_classes=7,\n        num_samples=my_num_samples,\n    ),\n])\n\n# validation dataset\nval_ds = CacheDataset(data=val_files, transform=base_transforms, cache_rate=1.0)\n\n# apply base transforms to validation dataset\nval_ds = Dataset(data=val_ds, transform=val_transforms)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=val_batch_size,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available(),\n    shuffle=False,\n)","metadata":{"papermill":{"duration":0.515869,"end_time":"2025-03-18T03:48:26.216236","exception":false,"start_time":"2025-03-18T03:48:25.700367","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:20.925719Z","iopub.execute_input":"2025-05-11T12:49:20.925935Z","iopub.status.idle":"2025-05-11T12:49:35.037934Z","shell.execute_reply.started":"2025-05-11T12:49:20.925916Z","shell.execute_reply":"2025-05-11T12:49:35.037077Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model setup","metadata":{"papermill":{"duration":0.013186,"end_time":"2025-03-18T03:48:26.243267","exception":false,"start_time":"2025-03-18T03:48:26.230081","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Copyright (c) MONAI Consortium\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#     http://www.apache.org/licenses/LICENSE-2.0\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom __future__ import annotations\n\nfrom collections.abc import Sequence\n\nimport torch\nimport torch.nn as nn\n\nfrom monai.networks.blocks.convolutions import Convolution\nfrom monai.networks.layers.factories import Norm\n\n__all__ = [\"AttentionUnet\"]\n\n\nclass ConvBlock(nn.Module):\n\n    def __init__(\n        self,\n        spatial_dims: int,\n        in_channels: int,\n        out_channels: int,\n        kernel_size: Sequence[int] | int = 3,\n        strides: int = 1,\n        dropout=0.0,\n    ):\n        super().__init__()\n        layers = [\n            Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=in_channels,\n                out_channels=out_channels,\n                kernel_size=kernel_size,\n                strides=strides,\n                padding=None,\n                adn_ordering=\"NDA\",\n                act=\"relu\",\n                norm=Norm.BATCH,\n                dropout=dropout,\n            ),\n            Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=out_channels,\n                out_channels=out_channels,\n                kernel_size=kernel_size,\n                strides=1,\n                padding=None,\n                adn_ordering=\"NDA\",\n                act=\"relu\",\n                norm=Norm.BATCH,\n                dropout=dropout,\n            ),\n        ]\n        self.conv = nn.Sequential(*layers)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_c: torch.Tensor = self.conv(x)\n        return x_c\n\n\nclass UpConv(nn.Module):\n\n    def __init__(self, spatial_dims: int, in_channels: int, out_channels: int, kernel_size=3, strides=2, dropout=0.0):\n        super().__init__()\n        self.up = Convolution(\n            spatial_dims,\n            in_channels,\n            out_channels,\n            strides=strides,\n            kernel_size=kernel_size,\n            act=\"relu\",\n            adn_ordering=\"NDA\",\n            norm=Norm.BATCH,\n            dropout=dropout,\n            is_transposed=True,\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_u: torch.Tensor = self.up(x)\n        return x_u\n\n\nclass AttentionBlock(nn.Module):\n\n    def __init__(self, spatial_dims: int, f_int: int, f_g: int, f_l: int, dropout=0.0):\n        super().__init__()\n        self.W_g = nn.Sequential(\n            Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=f_g,\n                out_channels=f_int,\n                kernel_size=1,\n                strides=1,\n                padding=0,\n                dropout=dropout,\n                conv_only=True,\n            ),\n            Norm[Norm.BATCH, spatial_dims](f_int),\n        )\n\n        self.W_x = nn.Sequential(\n            Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=f_l,\n                out_channels=f_int,\n                kernel_size=1,\n                strides=1,\n                padding=0,\n                dropout=dropout,\n                conv_only=True,\n            ),\n            Norm[Norm.BATCH, spatial_dims](f_int),\n        )\n\n        self.psi = nn.Sequential(\n            Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=f_int,\n                out_channels=1,\n                kernel_size=1,\n                strides=1,\n                padding=0,\n                dropout=dropout,\n                conv_only=True,\n            ),\n            Norm[Norm.BATCH, spatial_dims](1),\n            nn.Sigmoid(),\n        )\n\n        self.relu = nn.ReLU()\n\n    def forward(self, g: torch.Tensor, x: torch.Tensor) -> torch.Tensor:\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi: torch.Tensor = self.relu(g1 + x1)\n        psi = self.psi(psi)\n\n        return x * psi\n\n\nclass AttentionLayer(nn.Module):\n\n    def __init__(\n        self,\n        spatial_dims: int,\n        in_channels: int,\n        out_channels: int,\n        submodule: nn.Module,\n        up_kernel_size=3,\n        strides=2,\n        dropout=0.0,\n    ):\n        super().__init__()\n        self.attention = AttentionBlock(\n            spatial_dims=spatial_dims, f_g=in_channels, f_l=in_channels, f_int=in_channels // 2\n        )\n        self.upconv = UpConv(\n            spatial_dims=spatial_dims,\n            in_channels=out_channels,\n            out_channels=in_channels,\n            strides=strides,\n            kernel_size=up_kernel_size,\n        )\n        self.merge = Convolution(\n            spatial_dims=spatial_dims, in_channels=2 * in_channels, out_channels=in_channels, dropout=dropout\n        )\n        self.submodule = submodule\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        fromlower = self.upconv(self.submodule(x))\n        att = self.attention(g=fromlower, x=x)\n        att_m: torch.Tensor = self.merge(torch.cat((att, fromlower), dim=1))\n        return att_m\n\n\nclass AttentionUnet(nn.Module):\n    \"\"\"\n    Attention Unet based on\n    Otkay et al. \"Attention U-Net: Learning Where to Look for the Pancreas\"\n    https://arxiv.org/abs/1804.03999\n\n    Args:\n        spatial_dims: number of spatial dimensions of the input image.\n        in_channels: number of the input channel.\n        out_channels: number of the output classes.\n        channels (Sequence[int]): sequence of channels. Top block first. The length of `channels` should be no less than 2.\n        strides (Sequence[int]): stride to use for convolutions.\n        kernel_size: convolution kernel size.\n        up_kernel_size: convolution kernel size for transposed convolution layers.\n        dropout: dropout ratio. Defaults to no dropout.\n    \"\"\"\n\n    def __init__(\n        self,\n        spatial_dims: int,\n        in_channels: int,\n        out_channels: int,\n        channels: Sequence[int],\n        strides: Sequence[int],\n        kernel_size: Sequence[int] | int = 3,\n        up_kernel_size: Sequence[int] | int = 3,\n        dropout: float = 0.0,\n    ):\n        super().__init__()\n        self.dimensions = spatial_dims\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.channels = channels\n        self.strides = strides\n        self.kernel_size = kernel_size\n        self.dropout = dropout\n\n        head = ConvBlock(\n            spatial_dims=spatial_dims,\n            in_channels=in_channels,\n            out_channels=channels[0],\n            dropout=dropout,\n            kernel_size=self.kernel_size,\n        )\n        reduce_channels = Convolution(\n            spatial_dims=spatial_dims,\n            in_channels=channels[0],\n            out_channels=out_channels,\n            kernel_size=1,\n            strides=1,\n            padding=0,\n            conv_only=True,\n        )\n        self.up_kernel_size = up_kernel_size\n\n        def _create_block(channels: Sequence[int], strides: Sequence[int]) -> nn.Module:\n            if len(channels) > 2:\n                subblock = _create_block(channels[1:], strides[1:])\n                return AttentionLayer(\n                    spatial_dims=spatial_dims,\n                    in_channels=channels[0],\n                    out_channels=channels[1],\n                    submodule=nn.Sequential(\n                        ConvBlock(\n                            spatial_dims=spatial_dims,\n                            in_channels=channels[0],\n                            out_channels=channels[1],\n                            strides=strides[0],\n                            dropout=self.dropout,\n                            kernel_size=self.kernel_size,\n                        ),\n                        subblock,\n                    ),\n                    up_kernel_size=self.up_kernel_size,\n                    strides=strides[0],\n                    dropout=dropout,\n                )\n            else:\n                # the next layer is the bottom so stop recursion,\n                # create the bottom layer as the subblock for this layer\n                return self._get_bottom_layer(channels[0], channels[1], strides[0])\n\n        encdec = _create_block(self.channels, self.strides)\n        self.model = nn.Sequential(head, encdec, reduce_channels)\n\n    def _get_bottom_layer(self, in_channels: int, out_channels: int, strides: int) -> nn.Module:\n        return AttentionLayer(\n            spatial_dims=self.dimensions,\n            in_channels=in_channels,\n            out_channels=out_channels,\n            submodule=ConvBlock(\n                spatial_dims=self.dimensions,\n                in_channels=in_channels,\n                out_channels=out_channels,\n                strides=strides,\n                dropout=self.dropout,\n                kernel_size=self.kernel_size,\n            ),\n            up_kernel_size=self.up_kernel_size,\n            strides=strides,\n            dropout=self.dropout,\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_m: torch.Tensor = self.model(x)\n        return x_m","metadata":{"papermill":{"duration":0.034919,"end_time":"2025-03-18T03:48:26.291281","exception":false,"start_time":"2025-03-18T03:48:26.256362","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:35.039776Z","iopub.execute_input":"2025-05-11T12:49:35.040097Z","iopub.status.idle":"2025-05-11T12:49:35.064023Z","shell.execute_reply.started":"2025-05-11T12:49:35.040074Z","shell.execute_reply":"2025-05-11T12:49:35.063234Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n\nmodel = AttentionUnet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=7,\n    channels=(16, 32, 64, 128),\n    strides=(2, 2, 1),\n).to(device)\n\n# If multiple GPUs are available, wrap the model with DataParallel\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = nn.DataParallel(model)  # Wrap the model with DataParallel\n\nmodel = model.to(device)\n\nlr = 1e-3\noptimizer = torch.optim.Adam(model.parameters(), lr)\nloss_function = TverskyLoss(include_background=True, to_onehot_y=True, softmax=True)  # softmax=True for multiclass\ndice_metric = DiceMetric(include_background=False, reduction=\"mean\", ignore_empty=True) \nrecall_metric = ConfusionMatrixMetric(include_background=False, metric_name=\"recall\", reduction=\"None\")","metadata":{"papermill":{"duration":0.325536,"end_time":"2025-03-18T03:48:26.630443","exception":false,"start_time":"2025-03-18T03:48:26.304907","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:35.065116Z","iopub.execute_input":"2025-05-11T12:49:35.06543Z","iopub.status.idle":"2025-05-11T12:49:36.151494Z","shell.execute_reply.started":"2025-05-11T12:49:35.065398Z","shell.execute_reply":"2025-05-11T12:49:36.150619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"post_pred = AsDiscrete(argmax=True, to_onehot=len(root.pickable_objects)+1)\npost_label = AsDiscrete(to_onehot=len(root.pickable_objects)+1)\n\ndef train(train_loader, model, loss_function, metrics_function, optimizer, max_epochs=100):\n    val_interval = 2\n    best_metric = -1\n    best_metric_epoch = -1\n    epoch_loss_values = []\n    metric_values = []\n    for epoch in range(max_epochs):\n        print(\"-\" * 10)\n        print(f\"epoch {epoch + 1}/{max_epochs}\")\n        model.train()\n        epoch_loss = 0\n        step = 0\n        for batch_data in train_loader:\n            step += 1\n            inputs = batch_data[\"image\"].to(device)\n            labels = batch_data[\"label\"].to(device)\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = loss_function(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            epoch_loss += loss.item()\n            print(f\"batch {step}/{len(train_ds) // train_loader.batch_size}, \" f\"train_loss: {loss.item():.4f}\")\n        epoch_loss /= step\n        epoch_loss_values.append(epoch_loss)\n        print(f\"epoch {epoch + 1} average loss: {epoch_loss:.4f}\")\n        mlflow.log_metric(\"train_loss\", epoch_loss, step=epoch+1)\n\n        if (epoch + 1) % val_interval == 0:\n            model.eval()\n            with torch.no_grad():\n                for val_data in val_loader:\n                    val_inputs = val_data[\"image\"].to(device)\n                    val_labels = val_data[\"label\"].to(device)\n                    val_outputs = model(val_inputs)\n                    metric_val_outputs = [post_pred(i) for i in decollate_batch(val_outputs)]\n                    metric_val_labels = [post_label(i) for i in decollate_batch(val_labels)]\n                    \n                    \n                    # compute metric for current iteration\n                    metrics_function(y_pred=metric_val_outputs, y=metric_val_labels)\n\n                metrics = metrics_function.aggregate(reduction=\"mean_batch\")\n                metric_per_class = [\"{:.4g}\".format(x) for x in metrics]\n                metric = torch.mean(metrics).numpy(force=True)\n                mlflow.log_metric(\"validation metric\", metric, step=epoch+1)\n                for i,m in enumerate(metrics):\n                    mlflow.log_metric(f\"validation metric class {i+1}\", m, step=epoch+1)\n                metrics_function.reset()\n\n                metric_values.append(metric)\n                if metric > best_metric:\n                    best_metric = metric\n                    best_metric_epoch = epoch + 1\n                    torch.save(model.state_dict(), os.path.join('./', \"best_metric_model.pth\"))\n                    torch.save(model.state_dict(), os.path.join('./', \"best_model.bin\"))\n                    \n                    print(\"saved new best metric model\")\n                print(\n                    f\"current epoch: {epoch + 1} current mean recall per class: {', '.join(metric_per_class)}\"\n                    f\"\\nbest mean recall: {best_metric:.4f} \"\n                    f\"at epoch: {best_metric_epoch}\"\n                )","metadata":{"papermill":{"duration":0.025547,"end_time":"2025-03-18T03:48:26.669994","exception":false,"start_time":"2025-03-18T03:48:26.644447","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:36.152316Z","iopub.execute_input":"2025-05-11T12:49:36.152555Z","iopub.status.idle":"2025-05-11T12:49:36.163164Z","shell.execute_reply.started":"2025-05-11T12:49:36.152535Z","shell.execute_reply":"2025-05-11T12:49:36.162281Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training and tracking","metadata":{"papermill":{"duration":0.01306,"end_time":"2025-03-18T03:48:26.697327","exception":false,"start_time":"2025-03-18T03:48:26.684267","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from torchinfo import summary\n\nmlflow.end_run()\nmlflow.set_experiment('training 3D U-Net model for the cryoET ML Challenge')\nepochs = 200\nwith mlflow.start_run():\n    params = {\n        \"epochs\": epochs,\n        \"learning_rate\": lr,\n        \"loss_function\": loss_function.__class__.__name__,\n        \"metric_function\": recall_metric.__class__.__name__,\n        \"optimizer\": \"Adam\",\n    }\n    # Log training parameters.\n    mlflow.log_params(params)\n\n    # Log model summary.\n    with open(\"model_summary.txt\", \"w\") as f:\n        f.write(str(summary(model)))\n    mlflow.log_artifact(\"model_summary.txt\")\n\n    train(train_loader, model, loss_function, dice_metric, optimizer, max_epochs=epochs)\n\n    # Save the trained model to MLflow.\n    mlflow.pytorch.log_model(model, \"model\")\n\n","metadata":{"papermill":{"duration":3156.752094,"end_time":"2025-03-18T04:41:03.463063","exception":false,"start_time":"2025-03-18T03:48:26.710969","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:49:36.164009Z","iopub.execute_input":"2025-05-11T12:49:36.164241Z","execution_failed":"2025-05-11T13:01:55.182Z"}},"outputs":[],"execution_count":null}]}