{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook uses the code from the notebook at : https://github.com/czimaginginstitute/2024_czii_mlchallenge_notebooks/blob/main/3d_unet_monai/train.ipynb","metadata":{}},{"cell_type":"markdown","source":"## The following updates have been made:\n1. Visualizing the files and corresponding overlays in gif format to analyze the variations in tomography and overlay\n2. The model has been updated from UNet3D to to SegResNet3D.\n3. The code now uses distributed training to utilize Nvidia T4x2 provided by Kaggle.","metadata":{}},{"cell_type":"markdown","source":"### Comparison of SegResNet and UNet for Medical Image Segmentation\n\n| **Feature**             | **UNet**                                                                                     | **SegResNet**                                                                                 |\n|--------------------------|----------------------------------------------------------------------------------------------|----------------------------------------------------------------------------------------------|\n| **Architecture Type**   | Encoder-Decoder with symmetric skip connections                                              | Residual Network with deep residual blocks                                                  |\n| **Spatial Dimensions**  | Supports 2D and 3D input data                                                                | Supports 3D input data (SegResNet in MONAI is designed specifically for 3D)                 |\n| **Input Channels**      | Configurable (e.g., grayscale or multi-channel input)                                        | Configurable (e.g., grayscale or multi-channel input)                                        |\n| **Output Channels**     | Configurable (number of segmentation classes)                                                | Configurable (number of segmentation classes)                                                |\n| **Feature Extraction**  | Extracts features through convolutional layers and downsampling                              | Extracts features with deep residual blocks for better gradient flow                        |\n| **Skip Connections**    | Symmetric skip connections for fusing encoder and decoder features                           | Implicit residual connections within blocks for efficient gradient flow and learning        |\n| **Downsampling**        | Max pooling or strided convolutions                                                          | Strided convolutions within residual blocks                                                 |\n| **Upsampling**          | Transposed convolution (deconvolution)                                                       | Transposed convolution (deconvolution)                                                     |\n| **Initialization**      | Typically random initialization                                                              | `init_features` parameter determines the starting number of filters                         |\n| **Flexibility**         | Can easily be customized for specific tasks (e.g., depth of network, number of features)     | Less flexible but highly efficient for 3D medical image segmentation                        |\n| **Dropout**             | Optional, typically used in intermediate layers                                              | Configurable `dropout_prob` for regularization                                              |\n| **Use Cases**           | Widely used for 2D and 3D medical image segmentation tasks                                   | Best suited for 3D medical image segmentation tasks                                         |\n| **Advantages**          | - Easy to implement<br>- General-purpose architecture<br>- Effective for most segmentation tasks | - Better performance on complex 3D tasks<br>- Efficient gradient flow<br>- Deep residual learning |\n| **Disadvantages**       | - May require deeper architecture for highly complex tasks<br>- Susceptible to vanishing gradients | - Less flexible for 2D tasks<br>- Slightly more complex architecture                        |\n\n## Why SegResNet?\n\nSpecific requirements for our task which are satisfied by:\n- **SegResNet**: Specialized for 3D segmentation with improved gradient flow and efficiency due to residual connections, ideal for challenging 3D medical imaging tasks.\n- **SegResNet** for 3D tasks requiring deeper feature extraction and better gradient propagation.","metadata":{}},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Installing packages","metadata":{}},{"cell_type":"code","source":"!pip install zarr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:24:18.660686Z","iopub.execute_input":"2025-01-09T13:24:18.660969Z","iopub.status.idle":"2025-01-09T13:24:28.671470Z","shell.execute_reply.started":"2025-01-09T13:24:18.660945Z","shell.execute_reply":"2025-01-09T13:24:28.670659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install git+https://github.com/copick/copick-utils.git matplotlib tqdm copick \n!pip install -q \"monai-weekly[mlflow]\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:24:28.672679Z","iopub.execute_input":"2025-01-09T13:24:28.672945Z","iopub.status.idle":"2025-01-09T13:24:52.598722Z","shell.execute_reply.started":"2025-01-09T13:24:28.672922Z","shell.execute_reply":"2025-01-09T13:24:52.597903Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuring copick and loading the CryoET Tomography .zarr arrays","metadata":{}},{"cell_type":"code","source":"# Make a copick project\nimport 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-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            \"name\": \"membrane\",\n            \"is_particle\": false,\n            \"label\": 8,\n            \"color\": [100, 100, 100, 128]\n        },\n        {\n            \"name\": \"background\",\n            \"is_particle\": false,\n            \"label\": 9,\n            \"color\": [10, 150, 200, 128]\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": true\n    },\n\n    \"static_root\": \"/kaggle/input/czii-cryo-et-object-identification/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-cryo-et-object-identification/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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:25:38.823264Z","iopub.execute_input":"2025-01-09T13:25:38.823563Z","iopub.status.idle":"2025-01-09T13:25:39.115969Z","shell.execute_reply.started":"2025-01-09T13:25:38.823540Z","shell.execute_reply":"2025-01-09T13:25:39.115324Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Importing libraries","metadata":{}},{"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)\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:25:42.976202Z","iopub.execute_input":"2025-01-09T13:25:42.976509Z","iopub.status.idle":"2025-01-09T13:26:19.227679Z","shell.execute_reply.started":"2025-01-09T13:25:42.976480Z","shell.execute_reply":"2025-01-09T13:26:19.226789Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:26:19.228858Z","iopub.execute_input":"2025-01-09T13:26:19.229862Z","iopub.status.idle":"2025-01-09T13:26:19.234295Z","shell.execute_reply.started":"2025-01-09T13:26:19.229834Z","shell.execute_reply":"2025-01-09T13:26:19.233607Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Generate multi-class segmentation masks from picks, and saved them to the copick overlay directory (one-time)","metadata":{}},{"cell_type":"code","source":"from copick_utils.segmentation import segmentation_from_picks\nimport copick_utils.writers.write as write\nfrom collections import defaultdict\n\n# Just do this once\ngenerate_masks = True\n\nif generate_masks:\n    target_objects = defaultdict(dict)\n    for object in root.pickable_objects:\n        if object.is_particle:\n            target_objects[object.name]['label'] = object.label\n            target_objects[object.name]['radius'] = object.radius\n\n\n    for run in tqdm(root.runs):\n        tomo = run.get_voxel_spacing(10)\n        tomo = tomo.get_tomogram(tomo_type).numpy()\n        target = np.zeros(tomo.shape, dtype=np.uint8)\n        for pickable_object in root.pickable_objects:\n            pick = run.get_picks(object_name=pickable_object.name, user_id=\"curation\")\n            if len(pick):  \n                target = segmentation_from_picks.from_picks(pick[0], \n                                                            target, \n                                                            target_objects[pickable_object.name]['radius'] * 0.8,\n                                                            target_objects[pickable_object.name]['label']\n                                                            )\n        write.segmentation(run, target, copick_user_name, name=copick_segmentation_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:26:24.116390Z","iopub.execute_input":"2025-01-09T13:26:24.116715Z","iopub.status.idle":"2025-01-09T13:26:42.909671Z","shell.execute_reply.started":"2025-01-09T13:26:24.116686Z","shell.execute_reply":"2025-01-09T13:26:42.908797Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Get tomograms and their segmentaion masks (from picks) arrays\n","metadata":{}},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T14:26:23.697557Z","iopub.execute_input":"2025-01-09T14:26:23.697842Z","iopub.status.idle":"2025-01-09T14:26:30.687971Z","shell.execute_reply.started":"2025-01-09T14:26:23.697820Z","shell.execute_reply":"2025-01-09T14:26:30.686937Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualizing the files and corresponding overlays in gif format to analyze the variations in tomography and overlay","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom matplotlib.animation import FuncAnimation\n\n# Extract images and labels from data_dicts\nimages = data_dicts[0]['image']\nlabels = data_dicts[0]['label']\n\n# Create a figure for animation\nfig, axes = plt.subplots(1, 2, figsize=(15, 5))\n\n# Set up subplots\nax1, ax2 = axes\nim1 = ax1.imshow(images[0], cmap='gray', interpolation='none')\nax1.set_title(\"Tomogram\")\nax1.axis('off')\n\nim2 = ax2.imshow(labels[0], cmap='viridis', interpolation='none')\nax2.set_title(\"Painted Segmentation from Picks\")\nax2.axis('off')\n\n# Update function for animation\ndef update(frame):\n    im1.set_data(images[frame])\n    im2.set_data(labels[frame])\n    return [im1, im2]\n\n# Create the animation\nani = FuncAnimation(\n    fig, update, frames=len(images), interval=10000  # 1 frame per second\n)\n\n# Display the animation\nplt.close(fig)  # Close the static figure to display animation only\nfrom IPython.display import HTML\nHTML(ani.to_jshtml())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T14:30:23.187711Z","iopub.execute_input":"2025-01-09T14:30:23.188132Z","iopub.status.idle":"2025-01-09T14:30:32.605973Z","shell.execute_reply.started":"2025-01-09T14:30:23.188095Z","shell.execute_reply":"2025-01-09T14:30:32.604955Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Drag the slider for interactivity","metadata":{}},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n\n# # Plot the images\n# plt.figure(figsize=(15, 5))\n\n# plt.subplot(1, 2, 1)\n# plt.title('Tomogram')\n# plt.imshow(data_dicts[0]['image'][100],cmap='gray')\n# plt.axis('off')\n\n# plt.subplot(1, 2, 2)\n# plt.title('Painted Segmentation from Picks')\n# plt.imshow(data_dicts[0]['label'][100], cmap='viridis')\n# plt.axis('off')\n\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T14:29:10.377025Z","iopub.execute_input":"2025-01-09T14:29:10.377361Z","iopub.status.idle":"2025-01-09T14:29:10.381131Z","shell.execute_reply.started":"2025-01-09T14:29:10.377337Z","shell.execute_reply":"2025-01-09T14:29:10.380097Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Creating train and validation dataloaders","metadata":{}},{"cell_type":"code","source":"## using subset of the original dataset\n\nmy_num_samples = 16\ntrain_batch_size = 1\nval_batch_size = 1\n\ntrain_files, val_files = data_dicts[:5], data_dicts[5:7]\nprint(f\"Number of training samples: {len(train_files)}\")\nprint(f\"Number of validation samples: {len(val_files)}\")\n\n# Non-random transforms to be cached\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\n# Random transforms to be applied during training\nrandom_transforms = Compose([\n    RandCropByLabelClassesd(\n        keys=[\"image\", \"label\"],\n        label_key=\"label\",\n        spatial_size=[96, 96, 96],\n        num_classes=8,\n        num_samples=my_num_samples\n    ),\n    RandRotate90d(keys=[\"image\", \"label\"], prob=0.5, spatial_axes=[0, 2]),\n    RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=0),    \n])\n\n# Create the cached dataset with non-random transforms\ntrain_ds = CacheDataset(data=train_files, transform=non_random_transforms, cache_rate=1.0)\n\n# Wrap the cached dataset to apply random transforms during iteration\ntrain_ds = Dataset(data=train_ds, transform=random_transforms)\n\n# DataLoader remains the same\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    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    RandCropByLabelClassesd(\n        keys=[\"image\", \"label\"],\n        label_key=\"label\",\n        spatial_size=[96, 96, 96],\n        num_classes=8,\n        num_samples=my_num_samples,  # Use 1 to get a single, consistent crop per image\n    ),\n])\n\n# Create validation dataset\nval_ds = CacheDataset(data=val_files, transform=non_random_transforms, cache_rate=1.0)\n\n# Wrap the cached dataset to apply random transforms during iteration\nval_ds = Dataset(data=val_ds, transform=random_transforms)\n\n# Create validation DataLoader\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,  # Ensure the data order remains consistent\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:26:50.563939Z","iopub.execute_input":"2025-01-09T13:26:50.564321Z","iopub.status.idle":"2025-01-09T13:26:53.242464Z","shell.execute_reply.started":"2025-01-09T13:26:50.564285Z","shell.execute_reply":"2025-01-09T13:26:53.241579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install nnunet\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:26:53.243301Z","iopub.execute_input":"2025-01-09T13:26:53.243530Z","iopub.status.idle":"2025-01-09T13:26:53.247075Z","shell.execute_reply.started":"2025-01-09T13:26:53.243510Z","shell.execute_reply":"2025-01-09T13:26:53.246326Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Initiating the SegResNet model and configuring it for distributed training","metadata":{}},{"cell_type":"code","source":"import torch\nfrom monai.networks.nets import SegResNet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric, ConfusionMatrixMetric\n\n# Check device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Create a SegResNet model for multiclass segmentation\nmodel = SegResNet(\n    spatial_dims=3,  # Number of spatial dimensions (3 for 3D)\n    in_channels=1,  # Number of input channels (e.g., grayscale images)\n    out_channels=len(root.pickable_objects) + 1,  # Number of segmentation classes (including background)\n    # init_features=16,  # Initial number of features (filters in the first layer)\n    dropout_prob=0.2  # Dropout probability\n).to(device)\n\n# Define learning rate and optimizer\nlr = 1e-3\noptimizer = torch.optim.Adam(model.parameters(), lr)\n\n# Use TverskyLoss for multiclass segmentation\nloss_function = TverskyLoss(include_background=True, to_onehot_y=True, softmax=True)\n\n# Define the evaluation metrics\ndice_metric = DiceMetric(include_background=False, reduction=\"mean\", ignore_empty=True)\nrecall_metric = ConfusionMatrixMetric(include_background=False, metric_name=\"recall\", reduction=\"None\")\n\n# Example print statement for model summary\nprint(f\"SegResNet model: {model}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:29:40.094696Z","iopub.execute_input":"2025-01-09T13:29:40.095101Z","iopub.status.idle":"2025-01-09T13:29:40.471176Z","shell.execute_reply.started":"2025-01-09T13:29:40.095068Z","shell.execute_reply":"2025-01-09T13:29:40.470271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = torch.nn.DataParallel(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:29:46.330627Z","iopub.execute_input":"2025-01-09T13:29:46.330950Z","iopub.status.idle":"2025-01-09T13:29:46.335179Z","shell.execute_reply.started":"2025-01-09T13:29:46.330919Z","shell.execute_reply":"2025-01-09T13:29:46.334432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Defining the train loop","metadata":{}},{"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                    \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}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:29:49.184612Z","iopub.execute_input":"2025-01-09T13:29:49.184924Z","iopub.status.idle":"2025-01-09T13:29:49.194836Z","shell.execute_reply.started":"2025-01-09T13:29:49.184895Z","shell.execute_reply":"2025-01-09T13:29:49.194038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:29:50.728226Z","iopub.execute_input":"2025-01-09T13:29:50.728577Z","iopub.status.idle":"2025-01-09T13:29:50.732382Z","shell.execute_reply.started":"2025-01-09T13:29:50.728527Z","shell.execute_reply":"2025-01-09T13:29:50.731516Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training and validation","metadata":{}},{"cell_type":"code","source":"from torchinfo import summary\n\nmlflow.end_run()\nmlflow.set_experiment('training 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\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T13:29:51.867749Z","iopub.execute_input":"2025-01-09T13:29:51.868082Z","iopub.status.idle":"2025-01-09T14:17:36.705900Z","shell.execute_reply.started":"2025-01-09T13:29:51.868051Z","shell.execute_reply":"2025-01-09T14:17:36.704950Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## New modifications to improve the prediction will be available in the next version of this notebook. \n\nFollow me:\n* Kaggle: https://www.kaggle.com/akshat0007\n* Github: https://github.com/dubeyakshat07\n* Linkedin: https://www.linkedin.com/in/akshat0007/","metadata":{}}]}