{"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"},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":9869730,"sourceType":"datasetVersion","datasetId":6058495},{"sourceId":206640467,"sourceType":"kernelVersion"},{"sourceId":208227222,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# More Than Baseline UNet training + F beta metric + Optuna\n\n\nThis is built on top of the popular baseline code. \nMain Contribution:\n1. Incorporated the F beta metric into pytorch lightning pipeline. This metric is identical to the real metric for scoring\n2. Incorporated cca into the pipeline\n3. Wrapped the training process into Optuna trial\n\nIt's based on these three notebooks: \n\n1. [3D U-Net : Training Only](https://www.kaggle.com/code/ahsuna123/3d-u-net-training-only)\n2. [3D U-Net PyTorch Lightning distributed training](https://www.kaggle.com/code/zhuowenzhao11/3d-u-net-pytorch-lightning-distributed-training)\n3. [3d-unet using 2d image encoder](https://www.kaggle.com/code/hengck23/3d-unet-using-2d-image-encoder/notebook)\n4. [speed up connected component analysis with pytorch](https://www.kaggle.com/code/hengck23/speed-up-connected-component-analysis-with-pytorch)\n\nI've pre-computed the input data and stored them as numpy arrays so they don't have to be extracted every time the notebooks is run. ","metadata":{}},{"cell_type":"markdown","source":"## Installing offline deps\n\nAs this is a code comp, there is no internet. \nSo we have to do some silly things to get dependencies in here. \nWhy is asciitree such a PITA? ","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\n!rm -rf ./packages\ntry:\n    import zarr\nexcept: \n    !cp -r '/kaggle/input/hengck-czii-cryo-et-01/wheel_file' '/kaggle/working/'\n    !pip install /kaggle/working/wheel_file/asciitree-0.3.3/asciitree-0.3.3\n    !pip install --no-index --find-links=/kaggle/working/wheel_file zarr\n    !pip install --no-index --find-links=/kaggle/working/wheel_file connected-components-3d\nfrom typing import List, Tuple, Union\ndeps_path = '/kaggle/input/czii-cryoet-dependencies'\n! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt\nimport lightning.pytorch as pl\nfrom datetime import datetime\nimport pytz\nimport sys","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T09:20:26.593121Z","iopub.execute_input":"2025-01-26T09:20:26.594000Z","iopub.status.idle":"2025-01-26T09:20:50.160531Z","shell.execute_reply.started":"2025-01-26T09:20:26.593957Z","shell.execute_reply":"2025-01-26T09:20:50.159433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List, Tuple, Union\nimport numpy as np\nimport torch\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)","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:20:50.162810Z","iopub.execute_input":"2025-01-26T09:20:50.163479Z","iopub.status.idle":"2025-01-26T09:21:31.643934Z","shell.execute_reply.started":"2025-01-26T09:20:50.163441Z","shell.execute_reply":"2025-01-26T09:21:31.642665Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define some helper functions\n\n\n### Patching helper functions\n\nThese are mostly used to split large volumes into smaller ones and stitch them back together. ","metadata":{}},{"cell_type":"code","source":"def calculate_patch_starts(dimension_size: int, patch_size: int) -> List[int]:\n    \"\"\"\n    Calculate the starting positions of patches along a single dimension\n    with minimal overlap to cover the entire dimension.\n    \n    Parameters:\n    -----------\n    dimension_size : int\n        Size of the dimension\n    patch_size : int\n        Size of the patch in this dimension\n        \n    Returns:\n    --------\n    List[int]\n        List of starting positions for patches\n    \"\"\"\n    if dimension_size <= patch_size:\n        return [0]\n        \n    # Calculate number of patches needed\n    n_patches = np.ceil(dimension_size / patch_size)\n    \n    if n_patches == 1:\n        return [0]\n    \n    # Calculate overlap\n    total_overlap = (n_patches * patch_size - dimension_size) / (n_patches - 1)\n    \n    # Generate starting positions\n    positions = []\n    for i in range(int(n_patches)):\n        pos = int(i * (patch_size - total_overlap))\n        if pos + patch_size > dimension_size:\n            pos = dimension_size - patch_size\n        if pos not in positions:  # Avoid duplicates\n            positions.append(pos)\n    \n    return positions\n\ndef extract_3d_patches_minimal_overlap(arrays: List[np.ndarray], patch_size: int) -> Tuple[List[np.ndarray], List[Tuple[int, int, int]]]:\n    \"\"\"\n    Extract 3D patches from multiple arrays with minimal overlap to cover the entire array.\n    \n    Parameters:\n    -----------\n    arrays : List[np.ndarray]\n        List of input arrays, each with shape (m, n, l)\n    patch_size : int\n        Size of cubic patches (a x a x a)\n        \n    Returns:\n    --------\n    patches : List[np.ndarray]\n        List of all patches from all input arrays\n    coordinates : List[Tuple[int, int, int]]\n        List of starting coordinates (x, y, z) for each patch\n    \"\"\"\n    if not arrays or not isinstance(arrays, list):\n        raise ValueError(\"Input must be a non-empty list of arrays\")\n    \n    # Verify all arrays have the same shape\n    shape = arrays[0].shape\n    if not all(arr.shape == shape for arr in arrays):\n        raise ValueError(\"All input arrays must have the same shape\")\n    \n    if patch_size > min(shape):\n        raise ValueError(f\"patch_size ({patch_size}) must be smaller than smallest dimension {min(shape)}\")\n    \n    m, n, l = shape\n    patches = []\n    coordinates = []\n    \n    # Calculate starting positions for each dimension\n    x_starts = calculate_patch_starts(m, patch_size)\n    y_starts = calculate_patch_starts(n, patch_size)\n    z_starts = calculate_patch_starts(l, patch_size)\n    \n    # Extract patches from each array\n    for arr in arrays:\n        for x in x_starts:\n            for y in y_starts:\n                for z in z_starts:\n                    patch = arr[\n                        x:x + patch_size,\n                        y:y + patch_size,\n                        z:z + patch_size\n                    ]\n                    patches.append(patch)\n                    coordinates.append((x, y, z))\n    \n    return patches, coordinates\n\n# Note: I should probably averge the overlapping areas, \n# but here they are just overwritten by the most recent one. \n\ndef reconstruct_array(patches: List[np.ndarray], \n                     coordinates: List[Tuple[int, int, int]], \n                     original_shape: Tuple[int, int, int]) -> np.ndarray:\n    \"\"\"\n    Reconstruct array from patches.\n    \n    Parameters:\n    -----------\n    patches : List[np.ndarray]\n        List of patches to reconstruct from\n    coordinates : List[Tuple[int, int, int]]\n        Starting coordinates for each patch\n    original_shape : Tuple[int, int, int]\n        Shape of the original array\n        \n    Returns:\n    --------\n    np.ndarray\n        Reconstructed array\n    \"\"\"\n    reconstructed = np.zeros(original_shape, dtype=np.int64)  # To track overlapping regions\n    \n    patch_size = patches[0].shape[0]\n    \n    for patch, (x, y, z) in zip(patches, coordinates):\n        reconstructed[\n            x:x + patch_size,\n            y:y + patch_size,\n            z:z + patch_size\n        ] = patch\n        \n    \n    return reconstructed","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-01-26T09:21:31.645640Z","iopub.execute_input":"2025-01-26T09:21:31.646608Z","iopub.status.idle":"2025-01-26T09:21:31.661678Z","shell.execute_reply.started":"2025-01-26T09:21:31.646570Z","shell.execute_reply":"2025-01-26T09:21:31.660019Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission helper functions\n\nThese help with getting the submission in the correct format","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\nidx2label = {\n    1: 'apo-ferritin',\n    2: 'beta-amylase',\n    3: 'beta-galactosidase',\n    4: 'ribosome',\n    5: 'thyroglobulin',\n    6: 'virus-like-particle'\n}\n\ndef dict_to_df(coord_dict, experiment_name):    \n    # Create lists to store data\n    all_coords = []\n    all_labels = []\n    \n    # Process each label and its coordinates\n    for index, coords in coord_dict.items():\n        if index not in idx2label:\n            continue\n        label = idx2label[index]\n        all_coords.extend(coords)\n        all_labels.extend([label] * len(coords))\n    \n    # Concatenate all coordinates\n    all_coords = torch.vstack(all_coords)\n    print('stacked')\n    df = pd.DataFrame({\n        'experiment': experiment_name,\n        'particle_type': all_labels,\n        'x': all_coords[:, 0],\n        'y': all_coords[:, 1],\n        'z': all_coords[:, 2]\n    })\n    return df","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-01-26T09:21:31.664850Z","iopub.execute_input":"2025-01-26T09:21:31.665290Z","iopub.status.idle":"2025-01-26T09:21:31.684976Z","shell.execute_reply.started":"2025-01-26T09:21:31.665253Z","shell.execute_reply":"2025-01-26T09:21:31.683470Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# m = ConnectedComponentMetric(7, [0, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5])\n# i = 0\n# for b in valid_loader:\n#     y = b['label']\n#     y_squeezed = y.squeeze(1).long().cpu()\n#     y_squeezed[y_squeezed >= 7] = 0\n#     y_one_hot = torch.nn.functional.one_hot(y_squeezed, num_classes=7)\n#     y_one_hot = y_one_hot.permute(0, 4, 1, 2, 3)\n#     m.update(y_one_hot, y_one_hot)\n#     i = i + 1\n#     if i == 2:\n#         break\n# m.compute()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T09:21:31.686449Z","iopub.execute_input":"2025-01-26T09:21:31.686951Z","iopub.status.idle":"2025-01-26T09:21:31.703762Z","shell.execute_reply.started":"2025-01-26T09:21:31.686889Z","shell.execute_reply":"2025-01-26T09:21:31.702570Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reading in the data","metadata":{}},{"cell_type":"code","source":"TRAIN_DATA_DIR = \"/kaggle/input/create-numpy-dataset-exp-name\"\nTEST_DATA_DIR = \"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:31.705077Z","iopub.execute_input":"2025-01-26T09:21:31.705396Z","iopub.status.idle":"2025-01-26T09:21:31.716999Z","shell.execute_reply.started":"2025-01-26T09:21:31.705365Z","shell.execute_reply":"2025-01-26T09:21:31.715908Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_names = ['TS_5_4', 'TS_69_2', 'TS_6_6', 'TS_73_6', 'TS_86_3', 'TS_99_9']\nvalid_names = ['TS_6_4']\n\ntrain_files = []\nvalid_files = []\n\nfor name in train_names:\n    image = np.load(f\"{TRAIN_DATA_DIR}/train_image_{name}.npy\")\n    image_mean = np.mean(image)\n    image_std = np.std(image)\n    image = (image - image_mean) / image_std\n    label = np.load(f\"{TRAIN_DATA_DIR}/train_label_{name}.npy\")\n\n    train_files.append({\"image\": image, \"label\": label})\n    \n\nfor name in valid_names:\n    image = np.load(f\"{TRAIN_DATA_DIR}/train_image_{name}.npy\")\n    image_mean = np.mean(image)\n    image_std = np.std(image)\n    image = (image - image_mean) / image_std\n    label = np.load(f\"{TRAIN_DATA_DIR}/train_label_{name}.npy\")\n\n    valid_files.append({\"image\": image, \"label\": label})\n    \n","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:31.718137Z","iopub.execute_input":"2025-01-26T09:21:31.718469Z","iopub.status.idle":"2025-01-26T09:21:47.151848Z","shell.execute_reply.started":"2025-01-26T09:21:31.718431Z","shell.execute_reply":"2025-01-26T09:21:47.150720Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Functions to do connected component analysis","metadata":{}},{"cell_type":"code","source":"import cc3d\ndef np_find_connected_componet(probability, threshold):\n    probability = probability.numpy()\n    num_particle_type, D, H, W = probability.shape\n    binary = probability > np.array(threshold).reshape(num_particle_type,1,1,1)\n    componet = torch.zeros((num_particle_type, D, H, W), dtype=torch.int)\n    for i in range(num_particle_type):\n        cc = cc3d.connected_components(binary[i])\n        componet[i] = torch.tensor(cc, dtype=torch.int)\n    return componet\n\ndef np_find_centroid(component):\n    centroid =[]\n    component = component.numpy().astype(np.uint32)\n    num_particle_type, D, H, W = component.shape\n    for i in range(num_particle_type):\n        stats = cc3d.statistics(component[i])\n        zyx=stats['centroids'][1:]\n        xyz = np.ascontiguousarray(zyx[:,::-1].copy())\n        xyz = torch.tensor(xyz, dtype=torch.int)\n        centroid.append(xyz)\n    return centroid","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T09:21:47.153788Z","iopub.execute_input":"2025-01-26T09:21:47.154267Z","iopub.status.idle":"2025-01-26T09:21:47.184372Z","shell.execute_reply.started":"2025-01-26T09:21:47.154218Z","shell.execute_reply":"2025-01-26T09:21:47.183514Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Create the training dataloader\n\nI should probably find a way to create a dataloader that takes more batches. ","metadata":{}},{"cell_type":"code","source":"# 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\nraw_train_ds = CacheDataset(data=train_files, transform=non_random_transforms, cache_rate=1.0)","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:47.185620Z","iopub.execute_input":"2025-01-26T09:21:47.186510Z","iopub.status.idle":"2025-01-26T09:21:49.606719Z","shell.execute_reply.started":"2025-01-26T09:21:47.186473Z","shell.execute_reply":"2025-01-26T09:21:49.605536Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Create the validation dataloader\n\nHere I deviate a little from the source notebooks. \n\nIn the source, the validation dataloader also used the random transformations. This is bad practice and will result in noisy validation. \n\nHere I split the validation dataset in (slightly) overlapping blocks of `(96, 96 , 96)` so that we can have a consistent validation set that uses all the validation data. \n","metadata":{}},{"cell_type":"code","source":"val_images,val_labels = [dcts['image'] for dcts in valid_files],[dcts['label'] for dcts in valid_files]\n\nval_image_patches, _ = extract_3d_patches_minimal_overlap(val_images, 96)\nval_label_patches, _ = extract_3d_patches_minimal_overlap(val_labels, 96)\n\nval_patched_data = [{\"image\": img, \"label\": lbl} for img, lbl in zip(val_image_patches, val_label_patches)]\n\n\nvalid_ds = CacheDataset(data=val_patched_data, transform=non_random_transforms, cache_rate=1.0)\n\n\nvalid_batch_size = 16\n# DataLoader remains the same\nvalid_loader = DataLoader(\n    valid_ds,\n    batch_size=valid_batch_size,\n    shuffle=False,\n    num_workers=4,\n    drop_last=True,\n    pin_memory=torch.cuda.is_available()\n)","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:49.611286Z","iopub.execute_input":"2025-01-26T09:21:49.611971Z","iopub.status.idle":"2025-01-26T09:21:50.657171Z","shell.execute_reply.started":"2025-01-26T09:21:49.611931Z","shell.execute_reply":"2025-01-26T09:21:50.655928Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Initialize the model\n\nThis model is pretty much directly copied from [3D U-Net PyTorch Lightning distributed training](https://www.kaggle.com/code/zhuowenzhao11/3d-u-net-pytorch-lightning-distributed-training)","metadata":{}},{"cell_type":"code","source":"import lightning.pytorch as pl\n\nfrom monai.networks.nets import UNet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric\nfrom torchmetrics import Metric\n\nfrom metric import score\n\nclass ConnectedComponentMetric(Metric):\n    def __init__(self, num_classes, thresholds, dist_sync_on_step=False):\n        super().__init__(dist_sync_on_step=dist_sync_on_step)\n        self.num_classes = num_classes\n        self.thresholds = thresholds\n\n        # Add states to store results\n        self.add_state(\"predictions\", default=[], dist_reduce_fx=\"cat\")\n        self.add_state(\"labels\", default=[], dist_reduce_fx=\"cat\")\n\n    def update(self, y_hat, y):\n        # Accumulate predictions and ground truth\n        self.predictions.append(y_hat.cpu())\n        self.labels.append(y.cpu())\n\n    def compute(self):\n        # Concatenate predictions and labels across all batches\n        # Initialize dictionaries for storing centroids\n        pred_results = {}\n        label_results = {}\n        \n        all_data = torch.cat(self.predictions, dim=0)\n        # Process predictions\n        for prob in all_data:\n            # Find connected components for predictions\n            components = np_find_connected_componet(prob, self.thresholds)\n    \n            # Find centroids for connected components\n            pred_centroids = np_find_centroid(components)\n    \n            # Append results for each class to the prediction dictionary\n            for i, centroids in enumerate(pred_centroids):\n                if i not in pred_results:\n                    pred_results[i] = []\n                pred_results[i].extend(centroids)\n\n        \n        all_data = torch.cat(self.labels, dim=0)\n        # Process ground truth labels\n        for label in all_data:\n            # Find connected components for labels\n            components = np_find_connected_componet(label, self.thresholds)\n    \n            # Find centroids for connected components\n            label_centroids = np_find_centroid(components)\n    \n            # Append results for each class to the label dictionary\n            for i, centroids in enumerate(label_centroids):\n                if i not in label_results:\n                    label_results[i] = []\n                label_results[i].extend(centroids)\n\n        del all_data\n        # Convert the dictionaries into DataFrames\n        \n        pred_df = dict_to_df(pred_results, experiment_name=\"valid\")\n        del pred_results\n        print('pred df built')\n        label_df = dict_to_df(label_results, experiment_name=\"valid\")\n        del label_results\n        print('gt df built')\n        # Call the external scoring function with the two DataFrames\n        valid_score = score(\n            label_df,\n            pred_df,\n            'valid',\n            0.5,\n            4.0\n        )\n    \n        # Clear accumulated states (optional, for reuse in the next epoch)\n        self.reset()\n    \n        return valid_score\n        \nclass Model(pl.LightningModule):\n    def __init__(\n        self, \n        spatial_dims: int = 3,\n        in_channels: int = 1,\n        out_channels: int = 7,\n        channels: Union[Tuple[int, ...], List[int]] = (48, 64, 80, 80),\n        strides: Union[Tuple[int, ...], List[int]] = (2, 2, 1),\n        num_res_units: int = 1,\n        lr: float=1e-3):\n    \n        super().__init__()\n        self.save_hyperparameters()\n        self.model = UNet(\n            spatial_dims=self.hparams.spatial_dims,\n            in_channels=self.hparams.in_channels,\n            out_channels=self.hparams.out_channels,\n            channels=self.hparams.channels,\n            strides=self.hparams.strides,\n            num_res_units=self.hparams.num_res_units,\n        )\n        self.loss_fn = TverskyLoss(include_background=True, to_onehot_y=True, softmax=True)  # softmax=True for multiclass\n        # self.metric_fn = DiceMetric(include_background=False, reduction=\"mean\", ignore_empty=True)\n        self.metric_fn = ConnectedComponentMetric(num_classes=out_channels, thresholds=[-1, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5])\n\n        self.train_loss = 0\n        self.val_metric = 0\n        self.num_train_batch = 0\n        self.num_val_batch = 0\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch['image'], batch['label']\n        y_hat = self(x)\n        loss = self.loss_fn(y_hat, y)\n        self.train_loss += loss\n        self.num_train_batch += 1\n        torch.cuda.empty_cache()\n        return loss\n\n    def on_train_epoch_end(self):\n        loss_per_epoch = self.train_loss/self.num_train_batch\n        current_lr = self.trainer.optimizers[0].param_groups[0]['lr']\n        print(f\"Epoch {self.current_epoch} - Learning Rate: {current_lr:.6f} - Average Train Loss: {loss_per_epoch:.4f}\")\n        self.log('train_loss', loss_per_epoch)\n        self.train_loss = 0\n        self.num_train_batch = 0\n    \n    def validation_step(self, batch, batch_idx):\n        x, y = batch['image'], batch['label']\n        y_hat = self(x)\n\n        argmaxed = torch.argmax(y_hat, dim=1)\n        num_classes = y_hat.shape[1]\n        y_hat = torch.nn.functional.one_hot(argmaxed, num_classes=num_classes)\n        y_hat = y_hat.permute(0, 4, 1, 2, 3)\n\n        y_squeezed = y.squeeze(1).long().cpu()\n        y_squeezed[y_squeezed >= num_classes] = 0\n        y_one_hot = torch.nn.functional.one_hot(y_squeezed, num_classes=num_classes)\n        y_one_hot = y_one_hot.permute(0, 4, 1, 2, 3)\n        \n        # compute metric for current iteration\n        self.metric_fn.update(y_hat, y_one_hot)\n        self.num_val_batch += 1\n        return\n\n    def on_validation_epoch_end(self):\n        metric_per_epoch = self.metric_fn.compute()\n        print(f\"Epoch {self.current_epoch} - Average Val Metric: {metric_per_epoch:.4f}\")\n        self.log('val_metric', metric_per_epoch, sync_dist=False) # sync_dist=True for distributed training\n        self.val_metric = 0\n        self.num_val_batch = 0\n    \n    def configure_optimizers(self):\n        # Define the optimizer\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)\n\n        # Define the learning rate scheduler\n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer, \n            mode='max',  # Because we're maximizing the val_metric\n            factor=0.5,  # Reduce LR by 10x\n            patience=1,  # Number of epochs to wait before reducing LR\n            min_lr=1e-6  # Minimum learning rate\n        )\n\n        # Return both optimizer and scheduler\n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\n                \"scheduler\": scheduler,\n                \"monitor\": \"val_metric\",  # Monitors this metric to decide when to reduce LR\n                \"interval\": \"epoch\",     # Scheduler step frequency\n                \"frequency\": 1           # Step scheduler every epoch\n            }\n        }","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:50.659244Z","iopub.execute_input":"2025-01-26T09:21:50.659746Z","iopub.status.idle":"2025-01-26T09:21:50.705031Z","shell.execute_reply.started":"2025-01-26T09:21:50.659696Z","shell.execute_reply":"2025-01-26T09:21:50.703763Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Set up Optuna trials and Train\n\n","metadata":{}},{"cell_type":"code","source":"import optuna\nfrom lightning.pytorch import Trainer\nfrom lightning.pytorch.callbacks import EarlyStopping\nfrom lightning.pytorch.loggers import TensorBoardLogger\n\ndef permute_collate(batch):\n    all_images = []\n    all_labels = []\n\n    for item in batch:\n        # item[\"image\"].shape -> [num_samples, C, H, W, D]\n        # item[\"label\"].shape -> [num_samples, L, H, W, D]\n        for transformed in item:\n            all_images.append(transformed[\"image\"])\n            all_labels.append(transformed[\"label\"])\n\n    images_expanded = [img.unsqueeze(0) for img in all_images]  # shape -> [1, C, H, W, D]\n    labels_expanded = [lbl.unsqueeze(0) for lbl in all_labels] \n    images_flat = torch.cat(images_expanded, dim=0)\n    labels_flat = torch.cat(labels_expanded, dim=0)\n\n    # Generate a random permutation across the flattened dimension\n    permutation = torch.randperm(images_flat.size(0))\n    images_perm = images_flat[permutation]\n    labels_perm = labels_flat[permutation]\n\n    # Return a dictionary matching MONAI’s expected batch format\n    return {\"image\": images_perm, \"label\": labels_perm}\n    \ndef objective(trial):\n    torch.cuda.empty_cache()\n    # Hyperparameter suggestions\n    channels = trial.suggest_categorical(\"channels\", [(48, 64, 80, 96)])\n    strides = [2 for _ in range(len(channels)-1)] # trial.suggest_categorical(\"strides\", [(2, 2, 2), (2, 2, 1)])\n    num_res_units = trial.suggest_int(\"num_res_units\", 2, 4)\n    bs = trial.suggest_int(\"bs\", 1, 8, log=True)\n    lr = trial.suggest_float(\"lr\", 0.01, 0.03, log=True)\n\n    print(f\"Trial {trial.number} - Parameters: channels={channels}, strides={strides}, num_res_units={num_res_units}, lr={lr:.6f}, bs={bs}\")\n\n    my_num_samples = 2# 16 // bs\n    train_batch_size = bs\n    \n    # Random transforms to be applied during training\n    random_transforms = Compose([\n        RandCropByLabelClassesd(\n            keys=[\"image\", \"label\"],\n            label_key=\"label\",\n            spatial_size=[96, 96, 96],\n            num_classes=7,\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    train_ds = Dataset(data=raw_train_ds, transform=random_transforms)\n    \n    \n    # DataLoader remains the same\n    train_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        collate_fn=permute_collate\n    )\n\n    # Initialize the model\n    model = Model(channels=channels, strides=strides, num_res_units=num_res_units, lr=lr)\n\n    # TensorBoard Logger with unique name per trial\n    tb_logger = TensorBoardLogger(\n        save_dir=\"tensorboard_logs\"  # Base directory for TensorBoard logs\n        # name=f\"trial_{trial.number}\"  # Unique subdirectory for each trial\n    )\n    \n    # Trainer\n    trainer = Trainer(\n        max_epochs=15, \n        logger=tb_logger,\n        callbacks=[EarlyStopping(monitor=\"val_metric\", patience=3, mode=\"max\")],\n        precision=\"bf16-mixed\",  # Mixed precision with bfloat16\n        accelerator=\"gpu\",\n        devices=\"auto\"  # Automatically detect and use all available GPUs\n    )\n\n    # Fit the model\n    trainer.fit(model, train_loader, valid_loader)\n\n    # Log final validation metric\n    return trainer.callback_metrics[\"val_metric\"].item()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T09:21:50.706876Z","iopub.execute_input":"2025-01-26T09:21:50.707341Z","iopub.status.idle":"2025-01-26T09:21:50.722064Z","shell.execute_reply.started":"2025-01-26T09:21:50.707292Z","shell.execute_reply":"2025-01-26T09:21:50.720746Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Let there be gradients!\n\nLocally this config seems to train for about 1000 steps before the model starts overfitting. ","metadata":{}},{"cell_type":"code","source":"study = optuna.create_study(direction=\"maximize\")\nstudy.optimize(objective, n_trials=10)\n\nprint(f\"Best trial: {study.best_trial.value}\")\nprint(f\"Best hyperparameters: {study.best_trial.params}\")","metadata":{"execution":{"iopub.status.busy":"2025-01-26T09:21:50.723790Z","iopub.execute_input":"2025-01-26T09:21:50.724557Z","execution_failed":"2025-01-26T09:28:01.770Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Predict on the test set (Unchanged)\n\n","metadata":{}},{"cell_type":"code","source":"# model.eval();\n# model.to(\"cuda\");","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.770Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import json\n# copick_config_path = TRAIN_DATA_DIR + \"/copick.config\"\n\n# with open(copick_config_path) as f:\n#     copick_config = json.load(f)\n\n# copick_config['static_root'] = '/kaggle/input/czii-cryo-et-object-identification/test/static'\n\n# copick_test_config_path = 'copick_test.config'\n\n# with open(copick_test_config_path, 'w') as outfile:\n#     json.dump(copick_config, outfile)","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.770Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import copick\n\n# root = copick.from_file(copick_test_config_path)\n\n# copick_user_name = \"copickUtils\"\n# copick_segmentation_name = \"paintedPicks\"\n# voxel_size = 10\n# tomo_type = \"denoised\"","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Non-random transforms to be cached\n# inference_transforms = Compose([\n#     EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n#     NormalizeIntensityd(keys=\"image\"),\n#     Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n# ])","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import cc3d\n\n# id_to_name = {1: \"apo-ferritin\", \n#               2: \"beta-amylase\",\n#               3: \"beta-galactosidase\", \n#               4: \"ribosome\", \n#               5: \"thyroglobulin\", \n#               6: \"virus-like-particle\"}","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Iterate over test set\n\n\nBelow we will: \n1. Read in a run\n2. Split it into patches of size (96, 96, 96)\n3. Create a dataset from the patches\n4. Predict the segmentation mask\n5. Glue the mask back together\n6. Find the connected components for each class\n7. Find the centroids of the connected components\n8. Add to the dataframe\n\nThen do this for all runs. \n\nThis can probably be optimized quite a bit. ","metadata":{}},{"cell_type":"code","source":"# BLOB_THRESHOLD = 500\n# CERTAINTY_THRESHOLD = 0.5\n\n# classes = [1, 2, 3, 4, 5, 6]\n# with torch.no_grad():\n#     location_df = []\n#     for run in root.runs:\n#         print(run)\n\n#         tomo = run.get_voxel_spacing(10)\n#         tomo = tomo.get_tomogram(tomo_type).numpy()\n\n\n\n#         tomo_patches, coordinates  = extract_3d_patches_minimal_overlap([tomo], 96)\n\n#         tomo_patched_data = [{\"image\": img} for img in tomo_patches]\n\n#         tomo_ds = CacheDataset(data=tomo_patched_data, transform=inference_transforms, cache_rate=1.0)\n\n#         pred_masks = []\n\n#         for i in range(len(tomo_ds)):\n#             input_tensor = tomo_ds[i]['image'].unsqueeze(0).to(\"cuda\")\n#             model_output = model(input_tensor)\n\n#             probs = torch.softmax(model_output[0], dim=0)\n#             thresh_probs = probs > CERTAINTY_THRESHOLD\n#             _, max_classes = thresh_probs.max(dim=0)\n\n#             pred_masks.append(max_classes.cpu().numpy())\n            \n\n#         reconstructed_mask = reconstruct_array(pred_masks, coordinates, tomo.shape)\n        \n#         location = {}\n\n#         for c in classes:\n#             cc = cc3d.connected_components(reconstructed_mask == c)\n#             stats = cc3d.statistics(cc)\n#             zyx=stats['centroids'][1:]*10.012444 #https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895#3040071\n#             zyx_large = zyx[stats['voxel_counts'][1:] > BLOB_THRESHOLD]\n#             xyz =np.ascontiguousarray(zyx_large[:,::-1])\n\n#             location[id_to_name[c]] = xyz\n\n\n#         df = dict_to_df(location, run.name)\n#         location_df.append(df)\n    \n#     location_df = pd.concat(location_df)","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# location_df.insert(loc=0, column='id', value=np.arange(len(location_df)))\n# location_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !ls","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/hengck-czii-cryo-et-01/* .","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from czii_helper import *\n# from dataset import *\n# from scipy.optimize import linear_sum_assignment\n# import matplotlib.pyplot as plt","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os\n# if os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n#     MODE = 'submit'\n# else:\n#     MODE = 'local'\n\n\n\n\n\n\n\n# valid_dir ='/kaggle/input/czii-cryo-et-object-identification/train'\n# valid_id = ['TS_6_4', ]\n\n# def do_one_eval(truth, predict, threshold):\n#     P=len(predict)\n#     T=len(truth)\n\n#     if P==0:\n#         hit=[[],[]]\n#         miss=np.arange(T).tolist()\n#         fp=[]\n#         metric = [P,T,len(hit[0]),len(miss),len(fp)]\n#         return hit, fp, miss, metric\n\n#     if T==0:\n#         hit=[[],[]]\n#         fp=np.arange(P).tolist()\n#         miss=[]\n#         metric = [P,T,len(hit[0]),len(miss),len(fp)]\n#         return hit, fp, miss, metric\n\n#     #---\n#     distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n#     distance = distance**2\n#     distance = distance.sum(axis=2)\n#     distance = np.sqrt(distance)\n#     p_index, t_index = linear_sum_assignment(distance)\n\n#     valid = distance[p_index, t_index] <= threshold\n#     p_index = p_index[valid]\n#     t_index = t_index[valid]\n#     hit = [p_index.tolist(), t_index.tolist()]\n#     miss = np.arange(T)\n#     miss = miss[~np.isin(miss,t_index)].tolist()\n#     fp = np.arange(P)\n#     fp = fp[~np.isin(fp,p_index)].tolist()\n\n#     metric = [P,T,len(hit[0]),len(miss),len(fp)] #for lb metric F-beta copmutation\n#     return hit, fp, miss, metric\n\n\n# def compute_lb(submit_df, overlay_dir):\n#     valid_id = list(submit_df['experiment'].unique())\n#     print(valid_id)\n\n#     eval_df = []\n#     for id in valid_id:\n#         truth = read_one_truth(id, overlay_dir) #=f'{valid_dir}/overlay/ExperimentRuns')\n#         id_df = submit_df[submit_df['experiment'] == id]\n#         for p in PARTICLE:\n#             p = dotdict(p)\n#             print('\\r', id, p.name, end='', flush=True)\n#             xyz_truth = truth[p.name]\n#             xyz_predict = id_df[id_df['particle_type'] == p.name][['x', 'y', 'z']].values\n#             hit, fp, miss, metric = do_one_eval(xyz_truth, xyz_predict, p.radius* 0.5)\n#             eval_df.append(dotdict(\n#                 id=id, particle_type=p.name,\n#                 P=metric[0], T=metric[1], hit=metric[2], miss=metric[3], fp=metric[4],\n#             ))\n#     print('')\n#     eval_df = pd.DataFrame(eval_df)\n#     gb = eval_df.groupby('particle_type').agg('sum').drop(columns=['id'])\n#     gb.loc[:, 'precision'] = gb['hit'] / gb['P']\n#     gb.loc[:, 'precision'] = gb['precision'].fillna(0)\n#     gb.loc[:, 'recall'] = gb['hit'] / gb['T']\n#     gb.loc[:, 'recall'] = gb['recall'].fillna(0)\n#     gb.loc[:, 'f-beta4'] = 17 * gb['precision'] * gb['recall'] / (16 * gb['precision'] + gb['recall'])\n#     gb.loc[:, 'f-beta4'] = gb['f-beta4'].fillna(0)\n\n#     gb = gb.sort_values('particle_type').reset_index(drop=False)\n#     # https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895\n#     gb.loc[:, 'weight'] = [1, 0, 2, 1, 2, 1]\n#     lb_score = (gb['f-beta4'] * gb['weight']).sum() / gb['weight'].sum()\n#     return gb, lb_score\n\n\n# #debug\n# if 1:\n#     if MODE=='local':\n#     #if 1:\n#         submit_df=pd.read_csv(\n#            'submission.csv'\n#             # '/kaggle/input/hengck-czii-cryo-et-weights-01/submission.csv'\n#         )\n#         gb, lb_score = compute_lb(submit_df, f'{valid_dir}/overlay/ExperimentRuns')\n#         print(gb)\n#         print('lb_score:',lb_score)\n#         print('')\n\n\n#         #show one ----------------------------------\n#         fig = plt.figure(figsize=(18, 8))\n\n#         id = valid_id[0]\n#         truth = read_one_truth(id,overlay_dir=f'{valid_dir}/overlay/ExperimentRuns')\n\n#         submit_df = submit_df[submit_df['experiment']==id]\n#         for p in PARTICLE:\n#             p = dotdict(p)\n#             xyz_truth = truth[p.name]\n#             xyz_predict = submit_df[submit_df['particle_type']==p.name][['x','y','z']].values\n#             hit, fp, miss, _ = do_one_eval(xyz_truth, xyz_predict, p.radius)\n#             print(id, p.name)\n#             print('\\t num truth   :',len(xyz_truth) )\n#             print('\\t num predict :',len(xyz_predict) )\n#             print('\\t num hit  :',len(hit[0]) )\n#             print('\\t num fp   :',len(fp) )\n#             print('\\t num miss :',len(miss) )\n\n#             ax = fig.add_subplot(2, 3, p.label, projection='3d')\n#             if hit[0]:\n#                 pt = xyz_predict[hit[0]]\n#                 ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=0.5, color='r')\n#                 pt = xyz_truth[hit[1]]\n#                 ax.scatter(pt[:,0], pt[:,1], pt[:,2], s=80, facecolors='none', edgecolors='r')\n#             if fp:\n#                 pt = xyz_predict[fp]\n#                 ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], alpha=1, color='k')\n#             if miss:\n#                 pt = xyz_truth[miss]\n#                 ax.scatter(pt[:, 0], pt[:, 1], pt[:, 2], s=160, alpha=1, facecolors='none', edgecolors='k')\n\n#             ax.set_title(f'{p.name} ({p.difficulty})')\n\n#         plt.tight_layout()\n#         plt.show()\n        \n#         #--- \n#         zz=0","metadata":{"execution":{"execution_failed":"2025-01-26T09:28:01.771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}