{"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":"from monai.networks.blocks.convolutions import Convolution\nfrom monai.networks.layers.factories import Norm\nimport torch.nn as nn\nclass SingleConvolution(nn.Module):\n    \"\"\"\n    Auxiliary class to define a convolutional layer.\n    Each convolution block: convolution, batch normalization, ReLU activation.\n\n    Args:\n        nn.Module : receive the nn.Module properties\n    \"\"\"\n    def __init__(self, in_channels : int, out_channels : int, kernel_size: int = 3, padding: int = 1, stride : int = 1, bias: bool = False) -> None:\n        \"\"\"\n        Args:\n            in_channels (int): amount of input channels\n            out_channels (int): amount of output channels\n        \"\"\" \n        super(SingleConvolution, self).__init__()\n        \n        self.kernel_size = kernel_size\n        self.padding = padding\n        self.stride = stride\n        self.bias = bias\n        \n        self.singleConv = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size = self.kernel_size, padding = self.padding, stride = self.stride, bias = self.bias),\n            nn.BatchNorm3d(out_channels), \n            nn.ReLU(inplace=True),\n        )\n        \n    def forward(self, x : torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x (torch.Tensor): input tensor\n\n        Returns:\n            torch.Tensor: output tensor\n        \"\"\"\n        return self.singleConv(x)\n\nclass DoubleConvolution(nn.Module):\n    \"\"\"\n    Auxiliary class to define a convolutional layer.\n    Each convolution block: 3x3 convolution, batch normalization, ReLU activation.\n\n    Args:\n        nn.Module : receive the nn.Module properties\n    \"\"\"\n    def __init__(self, in_channels : int, out_channels : int, kernel_size: int = 3, padding: int = 1, stride: int = 1, bias: bool = False) -> None:\n        \"\"\"\n        Args:\n            in_channels (int): amount of input channels\n            out_channels (int): amount of output channels\n        \"\"\" \n        super(DoubleConvolution, self).__init__()\n        \n        self.conv1 = SingleConvolution(in_channels, out_channels, \n                                      kernel_size=kernel_size, \n                                      padding=padding, \n                                      stride=stride, \n                                      bias=bias)\n        self.conv2 = SingleConvolution(out_channels, out_channels, \n                                      kernel_size=kernel_size, \n                                      padding=padding, \n                                      stride=stride, \n                                      bias=bias)\n        \n    def forward(self, x : torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x (torch.Tensor): input tensor\n\n        Returns:\n            torch.Tensor: output tensor\n        \"\"\"\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return x\n\nclass ChannelGate(nn.Module): # C x 1 x 1\n    def __init__(self, in_channels: int, reduction_ratio: int = 16):\n        \n        super(ChannelGate, self).__init__()\n        \n        self.in_channels = in_channels\n        self.reduction_ratio = reduction_ratio\n        \n        self.squeezeAvg = nn.AdaptiveAvgPool3d(1) # Global Average Pooling\n        self.squeezeMax = nn.AdaptiveMaxPool3d(1) # Global Max Pooling\n        self.excitation = nn.Sequential(\n            nn.Linear(self.in_channels, self.in_channels // self.reduction_ratio),\n            nn.ReLU(inplace=True), # ReLU activation\n            nn.Linear(self.in_channels // self.reduction_ratio, self.in_channels),\n        )\n        self.sigActivation = nn.Sigmoid()\n        \n    def forward(self, x):\n        \n        batch_size, channels, _, _, _ = x.size()\n        \n        yAvg = self.squeezeAvg(x).view(batch_size, channels)\n        yAvg = self.excitation(yAvg).view(batch_size, channels, 1, 1, 1)\n        \n        yMax = self.squeezeMax(x).view(batch_size, channels)\n        yMax = self.excitation(yMax).view(batch_size, channels, 1, 1, 1)\n        \n        sum = yAvg + yMax\n        \n        return self.sigActivation(sum) * x\n\nclass SpatialGate(nn.Module): # 1 x H x W\n    def __init__(self, in_channels: int, kernel_size: int = 7, padding: int = 3, bias: bool = False):\n        \n        super(SpatialGate, self).__init__()\n        \n        self.kernel_size = kernel_size\n        self.in_channels = in_channels\n        self.padding = padding\n        self.bias = bias\n        \n        self.squeezeAvg = nn.AdaptiveAvgPool3d(1)\n        self.squeezeMax = nn.AdaptiveMaxPool3d(1)\n        self.spatial = SingleConvolution(2 * self.in_channels, self.in_channels, kernel_size = self.kernel_size, padding = self.padding)\n        self.sigActivation = nn.Sigmoid()\n        \n    def forward(self, x):\n        \n        yAvg = self.squeezeAvg(x)\n        yMax = self.squeezeMax(x)\n        y = torch.cat([yAvg, yMax], dim=1)\n        y = self.spatial(y)\n        \n        return self.sigActivation(y) * x\n\nclass CBAM(nn.Module):\n    def __init__(self, in_channels, reduction_ratio = 16, kernel_size: int = 7):\n        \n        super(CBAM, self).__init__()\n        \n        self.ChannelGate = ChannelGate(in_channels, reduction_ratio)\n        self.SpatialGate = SpatialGate(in_channels, kernel_size = kernel_size, padding = kernel_size // 2)\n        \n    def forward(self, x):\n        \n        x_out = self.ChannelGate(x)\n        x_out = self.SpatialGate(x_out)\n        \n        return x_out * x\n\nclass DownSampling(nn.Module):\n    \"\"\"\n    Auxiliary class to define a downsampling layer.\n    Each downsampling block: 2x2 max pooling, double convolution and squeeze and excitation.\n    input X output: [1, 16, 128, 128, 128] ->  [1, 32, 64, 64, 64] \n                    [1, 32, 64, 64, 64]    ->  [1, 64, 32, 32, 32]\n                    [1, 64, 32, 32, 32]    ->  [1, 128, 16, 16, 16]\n\n    Args:\n        nn.Module: receive the nn.Module properties\n    \"\"\"\n    def __init__(self, in_channels : int, out_channels : int, attention) -> None:\n        \"\"\"\n        Args:\n            in_channels (int): amount of input channels (16 or 32 or 64)\n            out_channels (int): amount of output channels (32 or 64 or 128)\n        \"\"\"        \n        super(DownSampling, self).__init__()\n        \n        self.attentionFunction = attention\n        \n        self.maxpool = nn.MaxPool3d(2)\n        self.conv = DoubleConvolution(in_channels, out_channels, kernel_size=3, padding=1, stride=1, bias=False)\n        self.attention = self.attentionFunction(out_channels)\n        \n    def forward(self, x : torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x (torch.Tensor): _description_\n\n        Returns:\n            torch.Tensor: _description_\n        \"\"\"\n        out = self.maxpool(x) # 2x2 max pooling -> 1/2 the size but same amount of channels\n        out = self.conv(out) # double convolution -> same size but double the amount of channels\n        out = self.attention(out) # squeeze and excitation\n        return out\n\nclass UpSampling(nn.Module):\n    \"\"\"\n    Auxiliary class to define a upsampling layer.\n    Each upsampling block: 2x2 upsampling, concatenation with skip connection, double convolution.\n    input X output: [1, 128, 16, 16, 16] ->  [1, 64, 32, 32, 32]\n                    [1, 64, 32, 32, 32]    ->  [1, 32, 64, 64, 64]\n                    [1, 32, 64, 64, 64]    ->  [1, 16, 128, 128, 128]\n                    \n    Args:\n        nn.Module: receive the nn.Module properties\n    \"\"\"\n    def __init__(self, in_channels: int, out_channels: int, bilinear: bool = False) -> None:\n        \"\"\"\n        Args:\n            in_channels (int): amount of input channels (128 or 64 or 32)\n            out_channels (int): amount of output channels (64 or 32 or 16)\n        \"\"\"\n        super(UpSampling, self).__init__()\n        \n        self.up = nn.ConvTranspose3d(in_channels, in_channels, kernel_size=2, stride=2)\n        self.conv = DoubleConvolution(int(in_channels + out_channels), out_channels, kernel_size=3, padding=1, stride=1, bias=False)\n        \n    def forward(self, x : torch.Tensor, skip_connection : torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x (torch.Tensor): the input tensor\n            skip_connection (torch.Tensor): the skip connection from the downsampling path\n\n        Returns:\n            torch.Tensor: the output tensor\n        \"\"\"\n        x = self.up(x) # 2x2 upsampling -> double the size but same amount of channels\n        x = torch.cat([skip_connection, x], dim=1) # concatenation with skip connection\n        out = self.conv(x) # double convolution -> same size but half the amount of channels\n        return out\n\nclass Attention_Unet(nn.Module):\n    def __init__(self, attentionFunction, in_channels=1, out_channels=1, channels=(32, 64, 128, 128)) -> None:\n        super(Attention_Unet, self).__init__()\n        self.attentionFunction = attentionFunction\n        layers = channels \n\n        self.attentionFunction = attentionFunction\n        \n        self.input = nn.Sequential(DoubleConvolution(in_channels, layers[0], kernel_size=3, padding = 1, stride=1, bias=False), self.attentionFunction(layers[0])) # tranform the input to 16 channels and apply squeeze and excitation\n        # encoding path\n        self.down1 = DownSampling(layers[0], layers[1], self.attentionFunction) \n        self.down2 = DownSampling(layers[1], layers[2], self.attentionFunction) \n        self.down3 = DownSampling(layers[2], layers[3], self.attentionFunction)\n        # decoding path\n        self.up1 = UpSampling(layers[3], layers[2])\n        self.up2 = UpSampling(layers[2], layers[1])\n        self.up3 = UpSampling(layers[1], layers[0])\n        self.output = nn.Sequential(nn.Conv3d(layers[0], out_channels, kernel_size=1)) # transform the output to 7 channel \n    \n    def forward(self, x : torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x (torch.Tensor): a tensor with shape [1, 1, 128, 128, 128]\n\n        Returns:\n            torch.Tensor: a tensor with shape [1, 1, 128, 128, 128]\n        \"\"\"\n        input = self.input(x) # [1, 1, 128, 128, 128] -> [1, 16, 128, 128, 128]\n        down1_output = self.down1(input)# [1, 16, 128, 128, 128] ->[1, 32, 64, 64, 64]\n        down2_output = self.down2(down1_output) # [1, 32, 64, 64, 64] -> [1, 64, 32, 32, 32]\n        down3_output = self.down3(down2_output) # [1, 64, 32, 32, 32] -> [1, 128, 16, 16, 16]\n        out = self.up1(down3_output, down2_output) # [1, 128, 16, 16, 16] -> [1, 64, 32, 32, 32]\n        out = self.up2(out, down1_output) # [1, 64, 32, 32, 32] -> [1, 32, 64, 64, 64]\n        out = self.up3(out, input) # [1, 32, 64, 64, 64] -> [1, 16, 128, 128, 128]\n        out = self.output(out) # [1, 16, 128, 128, 128] -> [1, 1, 128, 128, 128]\n        return out","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 = Attention_Unet(\n    attentionFunction=CBAM,\n    in_channels=1,\n    out_channels=7,\n    channels=(32,64,128,256),\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)  # must use onehot for multiclass\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}]}