{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install copick git+https://github.com/copick/copick-utils.git\n!pip install -q \"monai-weekly[mlflow]\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:46:09.768482Z","iopub.execute_input":"2024-12-23T22:46:09.768895Z","iopub.status.idle":"2024-12-23T22:46:23.110868Z","shell.execute_reply.started":"2024-12-23T22:46:09.768862Z","shell.execute_reply":"2024-12-23T22:46:23.109781Z"},"_kg_hide-output":false},"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-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\"\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)\n\nsource = '/kaggle/input/czii-cryo-et-object-identification/train/overlay'\ndest = '/kaggle/working/overlay'\n\nfor root, dirs, files in os.walk(source):\n    relpath = os.path.relpath(root, source)\n    target_dir = os.path.join(dest, relpath)\n    os.makedirs(target_dir, exist_ok=True)\n    \n    for file in files:\n        new_filename = f\"curation_0_{file}\"\n        source_file = os.path.join(root, file)\n        destination_file = os.path.join(target_dir, new_filename)\n        shutil.copy2(os.path.join(root, file), os.path.join(target_dir, new_filename))\n        print(f\"Copied {source_file} to {destination_file}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:46:23.112147Z","iopub.execute_input":"2024-12-23T22:46:23.112456Z","iopub.status.idle":"2024-12-23T22:46:23.195389Z","shell.execute_reply.started":"2024-12-23T22:46:23.112419Z","shell.execute_reply":"2024-12-23T22:46:23.194616Z"},"_kg_hide-output":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copick\nimport matplotlib.pyplot as plt\nimport torch\nimport torchinfo\nimport numpy as np\nimport zarr\nfrom tqdm import tqdm\nfrom copick_utils.segmentation.segmentation_from_picks import from_picks, downsample_to_exact_shape\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-12-23T22:46:23.196849Z","iopub.execute_input":"2024-12-23T22:46:23.197074Z","iopub.status.idle":"2024-12-23T22:46:23.202062Z","shell.execute_reply.started":"2024-12-23T22:46:23.197055Z","shell.execute_reply":"2024-12-23T22:46:23.201260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TOMO_TYPES = ['denoised', 'wbp']\nLEVELS = [\"0\"]\nVOXEL_SPACING = 10\nUSER_NAME = 'copickUtils'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:46:23.203334Z","iopub.execute_input":"2024-12-23T22:46:23.203625Z","iopub.status.idle":"2024-12-23T22:46:23.220361Z","shell.execute_reply.started":"2024-12-23T22:46:23.203596Z","shell.execute_reply":"2024-12-23T22:46:23.219613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Modified from copick_utils\ndef segmentation_from_picks(radius, painting_segmentation_name, run, voxel_spacing, tomo_type, pickable_object, pick_set, user_id=\"paintedPicks\", session_id=\"0\"):\n    # Fetch the tomogram and determine its multiscale structure\n    tomogram = run.get_voxel_spacing(voxel_spacing).get_tomograms(tomo_type)[0]\n    if not tomogram:\n        raise ValueError(\"Tomogram not found for the given parameters.\")\n\n    # Use copick to create a new segmentation if one does not exist\n    segs = run.get_segmentations(user_id=user_id, session_id=session_id, is_multilabel=True, name=painting_segmentation_name, voxel_size=voxel_spacing)\n    if len(segs) == 0:\n        seg = run.new_segmentation(voxel_spacing, painting_segmentation_name, session_id, True, user_id=user_id)\n    else:\n        seg = segs[0]\n\n    segmentation_group = zarr.open(seg.zarr(), mode=\"a\")\n    highest_res_name = \"0\"\n\n    # Get the highest resolution dimensions and create a new array if necessary\n    tomogram_zarr = zarr.open(tomogram.zarr(), \"r\")\n\n    highest_res_shape = tomogram_zarr[highest_res_name].shape\n    if highest_res_name not in segmentation_group:\n        segmentation_group.create(highest_res_name, shape=highest_res_shape, dtype=np.uint16, overwrite=True)\n\n    # Initialize or load the highest resolution array\n    highest_res_seg = segmentation_group[highest_res_name][:]\n\n    # Paint picks into the highest resolution array\n    highest_res_seg = from_picks(pick_set, highest_res_seg, radius, pickable_object.label, voxel_spacing)\n\n    # Write back the highest resolution data\n    segmentation_group[highest_res_name][:] = highest_res_seg\n\n    # Downsample to create lower resolution scales\n    multiscale_metadata = tomogram_zarr.attrs.get('multiscales', [{}])[0].get('datasets', [])\n    for level_index, level_metadata in enumerate(multiscale_metadata):\n        if level_index == 0:\n            continue\n\n        level_name = level_metadata.get(\"path\", str(level_index))\n        expected_shape = tuple(tomogram_zarr[level_name].shape)\n\n        # Compute scaling factors relative to the highest resolution shape\n        scaled_array = downsample_to_exact_shape(highest_res_seg, expected_shape)\n\n        # Create/overwrite the Zarr array for this level\n        segmentation_group.create_dataset(level_name, shape=expected_shape, data=scaled_array, dtype=np.uint16, overwrite=True)\n\n        segmentation_group[level_name][:] = scaled_array\n\n    return seg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:46:23.221379Z","iopub.execute_input":"2024-12-23T22:46:23.221687Z","iopub.status.idle":"2024-12-23T22:46:23.234889Z","shell.execute_reply.started":"2024-12-23T22:46:23.221636Z","shell.execute_reply":"2024-12-23T22:46:23.234190Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = []\n\nroot = copick.from_file(copick_config_path)\nfor run in tqdm(root.runs):\n    print(f\"\\nProcessing run: {run.meta.name}\")\n    for obj in root.pickable_objects:\n        if obj.is_particle:\n            print(f\"Processing {obj.name} with radius {obj.radius}\")\n            picks = run.get_picks(object_name=obj.name, user_id=\"curation\")\n            if picks:\n                seg = segmentation_from_picks(obj.radius, USER_NAME, run, VOXEL_SPACING, TOMO_TYPES[0], obj, picks[0])\n                print(f\"Created segmentation mask for {obj.name}\")\n    seg_group = zarr.open_group(seg.path)\n    for tomo_type in TOMO_TYPES:\n        tomogram = run.get_voxel_spacing(VOXEL_SPACING).get_tomograms(tomo_type)[0]\n        zarr_array = zarr.open(tomogram.zarr())\n        for level in LEVELS:\n            data.append({\"image\": np.array(zarr_array[level]), \"label\": np.array(seg_group[level], dtype=np.float32)})\n\nprint(len(data))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:46:23.235723Z","iopub.execute_input":"2024-12-23T22:46:23.235971Z","iopub.status.idle":"2024-12-23T22:47:24.197149Z","shell.execute_reply.started":"2024-12-23T22:46:23.235951Z","shell.execute_reply":"2024-12-23T22:47:24.196113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(15, 5))\nz = 100\n\nplt.subplot(1, 2, 1)\nplt.imshow(data[0]['image'][z], cmap='gray')\nplt.subplot(1, 2, 2)\nplt.imshow(data[0]['label'][z], cmap='viridis')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:47:24.198040Z","iopub.execute_input":"2024-12-23T22:47:24.198308Z","iopub.status.idle":"2024-12-23T22:47:24.696193Z","shell.execute_reply.started":"2024-12-23T22:47:24.198286Z","shell.execute_reply":"2024-12-23T22:47:24.695084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_samples = 16\ntrain_batch_size = 1\nval_batch_size = 1\n\ntrain_files, val_files = data[:10], [data[10], data[12]]\nprint(f\"Number of training samples: {len(train_files)}\")\nprint(f\"Number of validation samples: {len(val_files)}\")\n\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])\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=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])\ntrain_ds = CacheDataset(data=train_files, transform=non_random_transforms, cache_rate=1.0)\ntrain_ds = Dataset(data=train_ds, transform=random_transforms)\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\nval_ds = CacheDataset(data=val_files, transform=non_random_transforms, cache_rate=1.0)\nval_ds = Dataset(data=val_ds, transform=random_transforms)\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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:47:24.698041Z","iopub.execute_input":"2024-12-23T22:47:24.698263Z","iopub.status.idle":"2024-12-23T22:47:29.024009Z","shell.execute_reply.started":"2024-12-23T22:47:24.698243Z","shell.execute_reply":"2024-12-23T22:47:29.023201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 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=len(root.pickable_objects)+1,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=1,\n).to(device)\nlr = 1e-3\noptimizer = torch.optim.Adam(model.parameters(), lr)\nloss_function = TverskyLoss(include_background=True, to_onehot_y=True, softmax=True)\ndice_metric = DiceMetric(include_background=False, reduction=\"mean\", ignore_empty=True)\nrecall_metric = ConfusionMatrixMetric(include_background=False, metric_name=\"recall\", reduction=\"None\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:47:29.025362Z","iopub.execute_input":"2024-12-23T22:47:29.025742Z","iopub.status.idle":"2024-12-23T22:47:29.054241Z","shell.execute_reply.started":"2024-12-23T22:47:29.025706Z","shell.execute_reply":"2024-12-23T22:47:29.053468Z"}},"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}, 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        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                    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-12-23T22:47:29.054979Z","iopub.execute_input":"2024-12-23T22:47:29.055206Z","iopub.status.idle":"2024-12-23T22:47:29.064921Z","shell.execute_reply.started":"2024-12-23T22:47:29.055186Z","shell.execute_reply":"2024-12-23T22:47:29.063858Z"}},"outputs":[],"execution_count":null},{"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    mlflow.log_params(params)\n\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    mlflow.pytorch.log_model(model, \"model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:51:16.260288Z","iopub.execute_input":"2024-12-23T22:51:16.260710Z","iopub.status.idle":"2024-12-23T22:52:09.410573Z","shell.execute_reply.started":"2024-12-23T22:51:16.260673Z","shell.execute_reply":"2024-12-23T22:52:09.409873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = 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\nmodel.load_state_dict(torch.load(\"best_metric_model.pth\", weights_only=True))\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T22:52:16.313355Z","iopub.execute_input":"2024-12-23T22:52:16.313684Z","iopub.status.idle":"2024-12-23T22:52:16.359086Z","shell.execute_reply.started":"2024-12-23T22:52:16.313640Z","shell.execute_reply":"2024-12-23T22:52:16.358160Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    val_data = list(val_loader)[0]\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 = [AsDiscrete(argmax=True)(i) for i in decollate_batch(val_outputs)]\n    metric_val_labels = [AsDiscrete()(i) for i in decollate_batch(val_labels)]\n\nplt.figure(figsize=(15, 5))\nz = 50\n\nlabel = metric_val_labels[0].to(\"cpu\")\noutput = metric_val_outputs[0].to(\"cpu\")\n\nprint(label.shape, output.shape)\n\nplt.subplot(1, 2, 1)\nplt.imshow(label[0][z], cmap='viridis')\nplt.subplot(1, 2, 2)\nplt.imshow(output[0][z], cmap='viridis')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T23:03:56.093144Z","iopub.execute_input":"2024-12-23T23:03:56.093495Z","iopub.status.idle":"2024-12-23T23:04:00.120981Z","shell.execute_reply.started":"2024-12-23T23:03:56.093464Z","shell.execute_reply":"2024-12-23T23:04:00.119878Z"}},"outputs":[],"execution_count":null}]}