{"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# non random transforms\nnon_random_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\", \"label\"], axcodes=\"RAS\")\n])\n\nval_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# random transforms for training\nrandom_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    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 non random transforms to tarin dataset\ntrain_ds = CacheDataset(data=train_files, transform=non_random_transforms, cache_rate=1.0)\n\n# apply random transforms\ntrain_ds = Dataset(data=train_ds, transform=random_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_random_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=val_transforms, cache_rate=1.0)\n\n# apply random transforms to validation dataset\nval_ds = Dataset(data=val_ds, transform=val_random_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":"import torch.nn as nn\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n\nmodel = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=7,\n    channels=(16, 32, 64, 128),\n    strides=(2, 2, 1),\n    num_res_units=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)  # 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}]}