{"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":9869730,"sourceType":"datasetVersion","datasetId":6058495},{"sourceId":206165222,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Baseline UNet training + prediction/submission\n\n\nThis is the notebook I cobbled together to wrap my head around this challenge.\nI don't garuantee that the results are great, only that it works from end-to-end. \n\nIt trains a basic UNet and makes a submission. \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)\n\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":"deps_path = '/kaggle/input/czii-cryoet-dependencies'","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:34:51.067995Z","iopub.execute_input":"2024-11-24T13:34:51.068669Z","iopub.status.idle":"2024-11-24T13:34:51.076134Z","shell.execute_reply.started":"2024-11-24T13:34:51.068635Z","shell.execute_reply":"2024-11-24T13:34:51.075083Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! cp -r /kaggle/input/czii-cryoet-dependencies/asciitree-0.3.3/ asciitree-0.3.3/","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:34:51.359215Z","iopub.execute_input":"2024-11-24T13:34:51.359920Z","iopub.status.idle":"2024-11-24T13:34:52.494616Z","shell.execute_reply.started":"2024-11-24T13:34:51.359881Z","shell.execute_reply":"2024-11-24T13:34:52.493352Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip wheel asciitree-0.3.3/asciitree-0.3.3/\n","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-11-24T13:34:52.496636Z","iopub.execute_input":"2024-11-24T13:34:52.496967Z","iopub.status.idle":"2024-11-24T13:35:28.580444Z","shell.execute_reply.started":"2024-11-24T13:34:52.496934Z","shell.execute_reply":"2024-11-24T13:35:28.579517Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install asciitree-0.3.3-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:35:28.582018Z","iopub.execute_input":"2024-11-24T13:35:28.582320Z","iopub.status.idle":"2024-11-24T13:36:10.124880Z","shell.execute_reply.started":"2024-11-24T13:35:28.582288Z","shell.execute_reply":"2024-11-24T13:36:10.123770Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-11-24T13:36:10.127276Z","iopub.execute_input":"2024-11-24T13:36:10.127635Z","iopub.status.idle":"2024-11-24T13:36:30.277544Z","shell.execute_reply.started":"2024-11-24T13:36:10.127602Z","shell.execute_reply":"2024-11-24T13:36:30.276429Z"},"trusted":true},"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,\nScaleIntensityRangePercentilesd\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:37:47.737051Z","iopub.execute_input":"2024-11-24T13:37:47.737451Z","iopub.status.idle":"2024-11-24T13:37:47.742986Z","shell.execute_reply.started":"2024-11-24T13:37:47.737418Z","shell.execute_reply":"2024-11-24T13:37:47.741862Z"},"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":"2024-11-24T13:37:08.834824Z","iopub.execute_input":"2024-11-24T13:37:08.835559Z","iopub.status.idle":"2024-11-24T13:37:08.849191Z","shell.execute_reply.started":"2024-11-24T13:37:08.835497Z","shell.execute_reply":"2024-11-24T13:37:08.848193Z"},"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\ndef dict_to_df(coord_dict, experiment_name):\n    \"\"\"\n    Convert dictionary of coordinates to pandas DataFrame.\n    \n    Parameters:\n    -----------\n    coord_dict : dict\n        Dictionary where keys are labels and values are Nx3 coordinate arrays\n        \n    Returns:\n    --------\n    pd.DataFrame\n        DataFrame with columns ['x', 'y', 'z', 'label']\n    \"\"\"\n    # Create lists to store data\n    all_coords = []\n    all_labels = []\n    \n    # Process each label and its coordinates\n    for label, coords in coord_dict.items():\n        all_coords.append(coords)\n        all_labels.extend([label] * len(coords))\n    \n    # Concatenate all coordinates\n    all_coords = np.vstack(all_coords)\n    \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\n    \n    return df","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-11-24T13:37:08.850517Z","iopub.execute_input":"2024-11-24T13:37:08.850898Z","iopub.status.idle":"2024-11-24T13:37:08.878372Z","shell.execute_reply.started":"2024-11-24T13:37:08.850851Z","shell.execute_reply":"2024-11-24T13:37:08.877369Z"},"trusted":true},"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\"\nTEST_DATA_DIR = \"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:37:08.879416Z","iopub.execute_input":"2024-11-24T13:37:08.879942Z","iopub.status.idle":"2024-11-24T13:37:08.890855Z","shell.execute_reply.started":"2024-11-24T13:37:08.879914Z","shell.execute_reply":"2024-11-24T13:37:08.889940Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_list = []\nfor i in range(7):\n    image = np.load(f\"{TRAIN_DATA_DIR}/train_image_{i}.npy\")\n    label = np.load(f\"{TRAIN_DATA_DIR}/train_label_{i}.npy\")\n    print(np.unique(label),image.shape)\n    data_list.append({\"image\": image, \"label\": label})\n    \n\ntrain_files, val_files = data_list[:6], data_list[6:7]","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:37:08.891761Z","iopub.execute_input":"2024-11-24T13:37:08.892010Z","iopub.status.idle":"2024-11-24T13:37:32.400134Z","shell.execute_reply.started":"2024-11-24T13:37:08.891985Z","shell.execute_reply":"2024-11-24T13:37:32.399028Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import monai.transforms as mt\nimport numpy as np\nimport scipy.ndimage\n\nclass CopyPasteParticlesd(mt.MapTransform):\n    def __init__(self, keys, label_key, prob=0.5, max_copies=3):\n        super().__init__(keys)\n        self.label_key = label_key\n        self.prob = prob\n        self.max_copies = max_copies\n    \n    def __call__(self, data):\n        if np.random.rand() > self.prob:\n            return data\n\n        d = dict(data)\n        image, mask = d[\"image\"], d[self.label_key]\n\n        image = image.cpu().numpy()  # Convert torch to numpy\n        mask = mask.cpu().numpy()\n\n        mask = mask.squeeze(0)\n        image = image.squeeze(0)\n        \n        \n        # Find unique particles\n        for id in np.unique(mask)[1:]:  # Skip background\n\n            #We get a labeled mask for diff particles of same type\n            labeled_mask, num_features = scipy.ndimage.label(mask==id)\n            \n            if np.random.rand() > 0.5:\n\n                #select 0, to features+1\n                n_particles = np.random.randint(1, num_features + 1)\n                selected_particles = np.random.choice(range(1, num_features + 1), n_particles, replace=False)\n\n                for particle_id in selected_particles:\n                    \n                    where_particle = np.where(labeled_mask == particle_id)\n                    \n                    if len(where_particle[0]) == 0:\n                        continue\n                    \n                    z,y,x = where_particle\n                    \n                    min_z, max_z = np.min(z), np.max(z)\n                    min_y, max_y = np.min(y), np.max(y)\n                    min_x, max_x = np.min(x), np.max(x)\n                    \n                    # Extract particle\n                    particle_mask = labeled_mask[min_z:max_z+1, min_y:max_y+1, min_x:max_x+1] == particle_id\n                    particle_img = image[min_z:max_z+1, min_y:max_y+1, min_x:max_x+1] * particle_mask\n                    \n                    # Find new valid position\n                    for _ in range(20):  # Try 10 times to find valid position\n                        new_z = np.random.randint(0, labeled_mask.shape[0] - (max_z - min_z))\n                        new_y = np.random.randint(0, labeled_mask.shape[1] - (max_y - min_y))\n                        new_x = np.random.randint(0, labeled_mask.shape[2] - (max_x - min_x))\n                        \n                        # Check if position is empty\n                        target_region = mask[new_z:new_z+(max_z-min_z)+1, \n                                          new_y:new_y+(max_y-min_y)+1, \n                                          new_x:new_x+(max_x-min_x)+1]\n                        \n                        if np.sum(target_region) == 0:\n                            # Paste particle\n                            mask[new_z:new_z+(max_z-min_z)+1, new_y:new_y+(max_y-min_y)+1, new_x:new_x+(max_x-min_x)+1][particle_mask] = id\n                            image[new_z:new_z+(max_z-min_z)+1, new_y:new_y+(max_y-min_y)+1, new_x:new_x+(max_x-min_x)+1][particle_mask] = particle_img[particle_mask]\n                            break\n        \n        image = torch.from_numpy(image)  # Convert back to torch\n        mask = torch.from_numpy(mask)\n        \n        image = image.unsqueeze(0)\n        mask = mask.unsqueeze(0)\n\n        \n        d[\"image\"] = image\n        d[self.label_key] = mask\n        return d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-24T13:37:39.978556Z","iopub.execute_input":"2024-11-24T13:37:39.978920Z","iopub.status.idle":"2024-11-24T13:37:40.373849Z","shell.execute_reply.started":"2024-11-24T13:37:39.978887Z","shell.execute_reply":"2024-11-24T13:37:40.372862Z"}},"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(\n        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)\n\nmy_num_samples = 30\ntrain_batch_size = 1\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=6,\n        num_samples=my_num_samples\n    ),\n    CopyPasteParticlesd(keys=[\"image\",\"label\"],label_key=\"label\"),\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\ntrain_ds = Dataset(data=[d for d in raw_train_ds], transform=random_transforms)\n\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-24T21:49:01.145967Z","iopub.execute_input":"2024-11-24T21:49:01.146502Z","iopub.status.idle":"2024-11-24T21:49:01.365426Z","shell.execute_reply.started":"2024-11-24T21:49:01.146474Z","shell.execute_reply":"2024-11-24T21:49:01.364329Z"}},"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 val_files],[dcts['label'] for dcts in val_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\nvalid_ds = CacheDataset(data=val_patched_data, transform=non_random_transforms, cache_rate=1.0)\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    pin_memory=torch.cuda.is_available()\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-22T13:58:32.215375Z","iopub.execute_input":"2024-11-22T13:58:32.215680Z","iopub.status.idle":"2024-11-22T13:58:33.109817Z","shell.execute_reply.started":"2024-11-22T13:58:32.215650Z","shell.execute_reply":"2024-11-22T13:58:33.108931Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds[0]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-24T13:39:21.385770Z","iopub.execute_input":"2024-11-24T13:39:21.386472Z","iopub.status.idle":"2024-11-24T13:39:23.670534Z","shell.execute_reply.started":"2024-11-24T13:39:21.386431Z","shell.execute_reply":"2024-11-24T13:39:23.669563Z"}},"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\n\nclass Model(pl.LightningModule):\n    def __init__(\n        self, \n        spatial_dims: int = 3,\n        in_channels: int = 1,\n        out_channels: int = 6,\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\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        #print(f\"Epoch {self.current_epoch} - Average Train Loss: {loss_per_epoch:.4f}\")\n        self.log('train_loss', loss_per_epoch, prog_bar=True)\n        self.train_loss = 0\n        self.num_train_batch = 0\n    \n    def validation_step(self, batch, batch_idx):\n        with torch.no_grad(): # This ensures that gradients are not stored in memory\n            x, y = batch['image'], batch['label']\n            y_hat = self(x)\n            metric_val_outputs = [AsDiscrete(argmax=True, to_onehot=self.hparams.out_channels)(i) for i in decollate_batch(y_hat)]\n            metric_val_labels = [AsDiscrete(to_onehot=self.hparams.out_channels)(i) for i in decollate_batch(y)]\n\n            # compute metric for current iteration\n            self.metric_fn(y_pred=metric_val_outputs, y=metric_val_labels)\n            metrics = self.metric_fn.aggregate(reduction=\"mean_batch\")\n            val_metric = torch.mean(metrics) # I used mean over all particle species as the metric. This can be explored.\n            self.val_metric += val_metric \n            self.num_val_batch += 1\n        torch.cuda.empty_cache()\n        return {'val_metric': val_metric}\n\n    def on_validation_epoch_end(self):\n        metric_per_epoch = self.val_metric/self.num_val_batch\n        #print(f\"Epoch {self.current_epoch} - Average Val Metric: {metric_per_epoch:.4f}\")\n        self.log('val_metric', metric_per_epoch, prog_bar=True, 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        return torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:39:24.341674Z","iopub.execute_input":"2024-11-24T13:39:24.342565Z","iopub.status.idle":"2024-11-24T13:39:24.841331Z","shell.execute_reply.started":"2024-11-24T13:39:24.342489Z","shell.execute_reply":"2024-11-24T13:39:24.840546Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"channels = (48, 64, 80, 80)\nstrides_pattern = (2, 2, 1)       \nnum_res_units = 1\nlearning_rate = 1e-3\nnum_epochs = 50\n\nmodel = Model(channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)","metadata":{"execution":{"iopub.status.busy":"2024-11-24T13:39:26.664110Z","iopub.execute_input":"2024-11-24T13:39:26.664494Z","iopub.status.idle":"2024-11-24T13:39:26.694565Z","shell.execute_reply.started":"2024-11-24T13:39:26.664459Z","shell.execute_reply":"2024-11-24T13:39:26.693554Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train the model\n\n","metadata":{}},{"cell_type":"code","source":"torch.set_float32_matmul_precision('medium')\n\n# Check if CUDA is available and then count the GPUs\nif torch.cuda.is_available():\n    num_gpus = torch.cuda.device_count()\n    print(f\"Number of GPUs available: {num_gpus}\")\nelse:\n    print(\"No GPU available. Running on CPU.\")\ndevices = list(range(num_gpus))\nprint(devices)\n\n\ntrainer = pl.Trainer(\n    max_epochs=300,\n    #strategy=\"ddp_notebook\", \n    accelerator=\"gpu\",\n    devices=[0],# devices\n    num_nodes=1,\n    log_every_n_steps=10,\n    enable_progress_bar=True,\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-22T13:58:35.741942Z","iopub.execute_input":"2024-11-22T13:58:35.742205Z","iopub.status.idle":"2024-11-22T13:58:35.820599Z","shell.execute_reply.started":"2024-11-22T13:58:35.742169Z","shell.execute_reply":"2024-11-22T13:58:35.819852Z"},"trusted":true},"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":"\ntrainer.fit(model, train_loader, valid_loader)","metadata":{"execution":{"iopub.status.busy":"2024-11-22T13:58:35.821505Z","iopub.execute_input":"2024-11-22T13:58:35.821755Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Predict on the test set\n\n","metadata":{}},{"cell_type":"code","source":"model.eval();\nmodel.to(\"cuda\");","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\ncopick_config_path = TRAIN_DATA_DIR + \"/copick.config\"\n\nwith open(copick_config_path) as f:\n    copick_config = json.load(f)\n\ncopick_config['static_root'] = '/kaggle/input/czii-cryo-et-object-identification/test/static'\n\ncopick_test_config_path = 'copick_test.config'\n\nwith open(copick_test_config_path, 'w') as outfile:\n    json.dump(copick_config, outfile)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copick\n\nroot = copick.from_file(copick_test_config_path)\n\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Non-random transforms to be cached\ninference_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cc3d\n\nid_to_name = {1: \"apo-ferritin\", \n              2: \"beta-galactosidase\", \n              3: \"ribosome\", \n              4: \"thyroglobulin\", \n              5: \"virus-like-particle\"}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"BLOB_THRESHOLD = 300\nCERTAINTY_THRESHOLD = 0.5\n\nclasses = [1, 2, 3, 4, 5]\nwith 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":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"location_df.insert(loc=0, column='id', value=np.arange(len(location_df)))\nlocation_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}