{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Install pkgs","metadata":{}},{"cell_type":"markdown","source":"**Note:** This is training notebook only. Inference ain't included in . \nAnybody who wants to use this notebook for inference purposes is most welcome.","metadata":{}},{"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":"2024-11-07T10:38:30.308589Z","iopub.execute_input":"2024-11-07T10:38:30.308930Z","iopub.status.idle":"2024-11-07T10:39:35.018911Z","shell.execute_reply.started":"2024-11-07T10:38:30.308893Z","shell.execute_reply":"2024-11-07T10:39:35.017959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install zarr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-07T10:39:35.021085Z","iopub.execute_input":"2024-11-07T10:39:35.021436Z","iopub.status.idle":"2024-11-07T10:39:47.005401Z","shell.execute_reply.started":"2024-11-07T10:39:35.021402Z","shell.execute_reply":"2024-11-07T10:39:47.004339Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install copick","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-07T10:39:47.007201Z","iopub.execute_input":"2024-11-07T10:39:47.007654Z","iopub.status.idle":"2024-11-07T10:39:59.458776Z","shell.execute_reply.started":"2024-11-07T10:39:47.007609Z","shell.execute_reply":"2024-11-07T10:39:59.457747Z"}},"outputs":[],"execution_count":null},{"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":"2024-11-07T10:39:59.460450Z","iopub.execute_input":"2024-11-07T10:39:59.461128Z","iopub.status.idle":"2024-11-07T10:39:59.688771Z","shell.execute_reply.started":"2024-11-07T10:39:59.461080Z","shell.execute_reply":"2024-11-07T10:39:59.687859Z"}},"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)\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":"2024-11-07T10:39:59.691237Z","iopub.execute_input":"2024-11-07T10:39:59.691595Z","iopub.status.idle":"2024-11-07T10:40:47.608710Z","shell.execute_reply.started":"2024-11-07T10:39:59.691559Z","shell.execute_reply":"2024-11-07T10:40:47.607723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare the dataset\n## 1. Get copick root","metadata":{}},{"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":"2024-11-07T10:40:47.610079Z","iopub.execute_input":"2024-11-07T10:40:47.611018Z","iopub.status.idle":"2024-11-07T10:40:47.616968Z","shell.execute_reply.started":"2024-11-07T10:40:47.610980Z","shell.execute_reply":"2024-11-07T10:40:47.615922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. 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":"2024-11-07T10:40:47.618707Z","iopub.execute_input":"2024-11-07T10:40:47.619016Z","iopub.status.idle":"2024-11-07T10:41:03.076614Z","shell.execute_reply.started":"2024-11-07T10:40:47.618983Z","shell.execute_reply":"2024-11-07T10:41:03.075736Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Get tomograms and their segmentaion masks (from picks) arrays","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":"2024-11-07T10:41:03.077704Z","iopub.execute_input":"2024-11-07T10:41:03.077979Z","iopub.status.idle":"2024-11-07T10:41:09.753169Z","shell.execute_reply.started":"2024-11-07T10:41:03.077949Z","shell.execute_reply":"2024-11-07T10:41:09.752187Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Visualize the tomogram and painted segmentation from ground-truth picks","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Plot the images\nplt.figure(figsize=(15, 5))\n\nplt.subplot(1, 2, 1)\nplt.title('Tomogram')\nplt.imshow(data_dicts[0]['image'][100],cmap='gray')\nplt.axis('off')\n\nplt.subplot(1, 2, 2)\nplt.title('Painted Segmentation from Picks')\nplt.imshow(data_dicts[0]['label'][100], cmap='viridis')\nplt.axis('off')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-07T10:41:09.754463Z","iopub.execute_input":"2024-11-07T10:41:09.755133Z","iopub.status.idle":"2024-11-07T10:41:10.230095Z","shell.execute_reply.started":"2024-11-07T10:41:09.755085Z","shell.execute_reply":"2024-11-07T10:41:10.229174Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Prepare dataloaders","metadata":{}},{"cell_type":"code","source":"my_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":"2024-11-07T10:41:10.231223Z","iopub.execute_input":"2024-11-07T10:41:10.231580Z","iopub.status.idle":"2024-11-07T10:41:12.385262Z","shell.execute_reply.started":"2024-11-07T10:41:10.231546Z","shell.execute_reply":"2024-11-07T10:41:12.384114Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model setup","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n# Create UNet, DiceLoss and Adam optimizer\nmodel = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=len(root.pickable_objects)+1,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=1,\n).to(device)\n\nlr = 1e-3\noptimizer = torch.optim.Adam(model.parameters(), lr)\n#loss_function = DiceLoss(include_background=True, to_onehot_y=True, softmax=True)  # softmax=True for multiclass\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":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-07T10:41:12.386887Z","iopub.execute_input":"2024-11-07T10:41:12.387554Z","iopub.status.idle":"2024-11-07T10:41:12.602210Z","shell.execute_reply.started":"2024-11-07T10:41:12.387518Z","shell.execute_reply":"2024-11-07T10:41:12.601467Z"}},"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=25):\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}\"\n                )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-07T10:41:12.603782Z","iopub.execute_input":"2024-11-07T10:41:12.604357Z","iopub.status.idle":"2024-11-07T10:41:12.619481Z","shell.execute_reply.started":"2024-11-07T10:41:12.604306Z","shell.execute_reply":"2024-11-07T10:41:12.618549Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training and tracking","metadata":{}},{"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 = 50\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":{"execution":{"iopub.status.busy":"2024-11-07T10:54:27.703840Z","iopub.execute_input":"2024-11-07T10:54:27.704330Z","iopub.status.idle":"2024-11-07T11:04:03.078262Z","shell.execute_reply.started":"2024-11-07T10:54:27.704285Z","shell.execute_reply":"2024-11-07T11:04:03.077243Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Note:** Code taken from the official CZ Imaging Institute's official github page.","metadata":{}}]}