{"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"}],"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-19T21:09:38.661767Z","iopub.execute_input":"2024-11-19T21:09:38.662016Z","iopub.status.idle":"2024-11-19T21:09:38.669631Z","shell.execute_reply.started":"2024-11-19T21:09:38.661990Z","shell.execute_reply":"2024-11-19T21:09:38.668736Z"},"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-19T21:09:38.670965Z","iopub.execute_input":"2024-11-19T21:09:38.671222Z","iopub.status.idle":"2024-11-19T21:09:39.764929Z","shell.execute_reply.started":"2024-11-19T21:09:38.671196Z","shell.execute_reply":"2024-11-19T21:09:39.762108Z"},"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-19T21:09:39.768889Z","iopub.execute_input":"2024-11-19T21:09:39.770750Z","iopub.status.idle":"2024-11-19T21:10:15.808402Z","shell.execute_reply.started":"2024-11-19T21:09:39.769291Z","shell.execute_reply":"2024-11-19T21:10:15.807299Z"},"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-19T21:10:15.809846Z","iopub.execute_input":"2024-11-19T21:10:15.810219Z","iopub.status.idle":"2024-11-19T21:10:56.396466Z","shell.execute_reply.started":"2024-11-19T21:10:15.810182Z","shell.execute_reply":"2024-11-19T21:10:56.395392Z"},"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-19T21:10:56.399282Z","iopub.execute_input":"2024-11-19T21:10:56.400116Z","iopub.status.idle":"2024-11-19T21:11:14.619378Z","shell.execute_reply.started":"2024-11-19T21:10:56.400081Z","shell.execute_reply":"2024-11-19T21:11:14.618544Z"},"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,\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:11:14.620928Z","iopub.execute_input":"2024-11-19T21:11:14.621351Z","iopub.status.idle":"2024-11-19T21:11:47.939389Z","shell.execute_reply.started":"2024-11-19T21:11:14.621294Z","shell.execute_reply":"2024-11-19T21:11:47.938583Z"},"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-19T21:11:47.940809Z","iopub.execute_input":"2024-11-19T21:11:47.941727Z","iopub.status.idle":"2024-11-19T21:11:47.956594Z","shell.execute_reply.started":"2024-11-19T21:11:47.941688Z","shell.execute_reply":"2024-11-19T21:11:47.955741Z"},"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-19T21:11:47.957605Z","iopub.execute_input":"2024-11-19T21:11:47.957870Z","iopub.status.idle":"2024-11-19T21:11:47.978432Z","shell.execute_reply.started":"2024-11-19T21:11:47.957844Z","shell.execute_reply":"2024-11-19T21:11:47.977741Z"},"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-exp-name\"\nTEST_DATA_DIR = \"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:11:47.979396Z","iopub.execute_input":"2024-11-19T21:11:47.979657Z","iopub.status.idle":"2024-11-19T21:11:47.989983Z","shell.execute_reply.started":"2024-11-19T21:11:47.979632Z","shell.execute_reply":"2024-11-19T21:11:47.989174Z"},"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    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    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":"2024-11-19T21:11:47.990874Z","iopub.execute_input":"2024-11-19T21:11:47.991470Z","iopub.status.idle":"2024-11-19T21:12:03.216903Z","shell.execute_reply.started":"2024-11-19T21:11:47.991441Z","shell.execute_reply":"2024-11-19T21:12:03.216155Z"},"trusted":true},"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)\n\n\nmy_num_samples = 16\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=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\ntrain_ds = Dataset(data=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":{"execution":{"iopub.status.busy":"2024-11-19T21:12:03.220682Z","iopub.execute_input":"2024-11-19T21:12:03.220954Z","iopub.status.idle":"2024-11-19T21:12:05.365905Z","shell.execute_reply.started":"2024-11-19T21:12:03.220928Z","shell.execute_reply":"2024-11-19T21:12:05.364898Z"},"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    pin_memory=torch.cuda.is_available()\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:12:05.366903Z","iopub.execute_input":"2024-11-19T21:12:05.367204Z","iopub.status.idle":"2024-11-19T21:12:06.241225Z","shell.execute_reply.started":"2024-11-19T21:12:05.367172Z","shell.execute_reply":"2024-11-19T21:12:06.240182Z"},"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\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\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-19T21:12:06.242509Z","iopub.execute_input":"2024-11-19T21:12:06.242797Z","iopub.status.idle":"2024-11-19T21:12:07.061144Z","shell.execute_reply.started":"2024-11-19T21:12:06.242769Z","shell.execute_reply":"2024-11-19T21:12:07.060486Z"},"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 = 100\n\nmodel = Model(channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:12:07.062162Z","iopub.execute_input":"2024-11-19T21:12:07.062446Z","iopub.status.idle":"2024-11-19T21:12:07.089742Z","shell.execute_reply.started":"2024-11-19T21:12:07.062419Z","shell.execute_reply":"2024-11-19T21:12:07.089130Z"},"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=num_epochs,\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-19T21:12:07.090721Z","iopub.execute_input":"2024-11-19T21:12:07.091063Z","iopub.status.idle":"2024-11-19T21:12:07.172294Z","shell.execute_reply.started":"2024-11-19T21:12:07.091022Z","shell.execute_reply":"2024-11-19T21:12:07.171583Z"},"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-19T21:12:07.173285Z","iopub.execute_input":"2024-11-19T21:12:07.173669Z","iopub.status.idle":"2024-11-19T21:32:20.767852Z","shell.execute_reply.started":"2024-11-19T21:12:07.173630Z","shell.execute_reply":"2024-11-19T21:32:20.767015Z"},"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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:20.769598Z","iopub.execute_input":"2024-11-19T21:32:20.769975Z","iopub.status.idle":"2024-11-19T21:32:20.779771Z","shell.execute_reply.started":"2024-11-19T21:32:20.769933Z","shell.execute_reply":"2024-11-19T21:32:20.779052Z"},"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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:20.781160Z","iopub.execute_input":"2024-11-19T21:32:20.781430Z","iopub.status.idle":"2024-11-19T21:32:20.820128Z","shell.execute_reply.started":"2024-11-19T21:32:20.781402Z","shell.execute_reply":"2024-11-19T21:32:20.819279Z"},"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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:20.821406Z","iopub.execute_input":"2024-11-19T21:32:20.822030Z","iopub.status.idle":"2024-11-19T21:32:21.592482Z","shell.execute_reply.started":"2024-11-19T21:32:20.821982Z","shell.execute_reply":"2024-11-19T21:32:21.591544Z"},"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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:21.593755Z","iopub.execute_input":"2024-11-19T21:32:21.595036Z","iopub.status.idle":"2024-11-19T21:32:21.600315Z","shell.execute_reply.started":"2024-11-19T21:32:21.594993Z","shell.execute_reply":"2024-11-19T21:32:21.599493Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cc3d\n\nid_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":{"iopub.status.busy":"2024-11-19T21:32:21.601645Z","iopub.execute_input":"2024-11-19T21:32:21.601913Z","iopub.status.idle":"2024-11-19T21:32:21.619239Z","shell.execute_reply.started":"2024-11-19T21:32:21.601886Z","shell.execute_reply":"2024-11-19T21:32:21.618480Z"},"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\nCERTAINTY_THRESHOLD = 0.5\n\nclasses = [1, 2, 3, 4, 5, 6]\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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:21.620215Z","iopub.execute_input":"2024-11-19T21:32:21.620928Z","iopub.status.idle":"2024-11-19T21:32:57.567719Z","shell.execute_reply.started":"2024-11-19T21:32:21.620899Z","shell.execute_reply":"2024-11-19T21:32:57.566989Z"},"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":{"execution":{"iopub.status.busy":"2024-11-19T21:32:57.568926Z","iopub.execute_input":"2024-11-19T21:32:57.569560Z","iopub.status.idle":"2024-11-19T21:32:57.586987Z","shell.execute_reply.started":"2024-11-19T21:32:57.569519Z","shell.execute_reply":"2024-11-19T21:32:57.586341Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:32:57.588065Z","iopub.execute_input":"2024-11-19T21:32:57.588622Z","iopub.status.idle":"2024-11-19T21:32:58.767534Z","shell.execute_reply.started":"2024-11-19T21:32:57.588582Z","shell.execute_reply":"2024-11-19T21:32:58.766498Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/hengck-czii-cryo-et-01/* .","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:32:58.768940Z","iopub.execute_input":"2024-11-19T21:32:58.769243Z","iopub.status.idle":"2024-11-19T21:33:00.537355Z","shell.execute_reply.started":"2024-11-19T21:32:58.769214Z","shell.execute_reply":"2024-11-19T21:33:00.536280Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from czii_helper import *\nfrom dataset import *\nfrom scipy.optimize import linear_sum_assignment\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2024-11-19T21:33:00.539136Z","iopub.execute_input":"2024-11-19T21:33:00.540034Z","iopub.status.idle":"2024-11-19T21:33:00.547260Z","shell.execute_reply.started":"2024-11-19T21:33:00.539982Z","shell.execute_reply":"2024-11-19T21:33:00.546595Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    MODE = 'submit'\nelse:\n    MODE = 'local'\n\n\n\n\n\n\n\nvalid_dir ='/kaggle/input/czii-cryo-et-object-identification/train'\nvalid_id = ['TS_6_4', ]\n\ndef 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\ndef 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\nif 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":{"iopub.status.busy":"2024-11-19T21:49:29.234206Z","iopub.execute_input":"2024-11-19T21:49:29.234690Z","iopub.status.idle":"2024-11-19T21:49:30.805617Z","shell.execute_reply.started":"2024-11-19T21:49:29.234643Z","shell.execute_reply":"2024-11-19T21:49:30.804545Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}