{"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":10011619,"sourceType":"datasetVersion","datasetId":6144035},{"sourceId":208775456,"sourceType":"kernelVersion"},{"sourceId":210123055,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"mode=\"SUBMIT\"\nif mode == \"SUBMIT\":\n    deps_path = '/kaggle/input/czii-cryoet-dependencies'\n    ! cp -r /kaggle/input/czii-cryoet-dependencies/asciitree-0.3.3/ asciitree-0.3.3/\n    ! pip wheel asciitree-0.3.3/asciitree-0.3.3/\n    ! pip install asciitree-0.3.3-py3-none-any.whl\n    ! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt\nelse:\n    !pip install monai\n    !pip install lightning\n    !pip install connected-components-3d copick\n    !pip install pickle\n    !pip install zarr\n    ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:11:58.196982Z","iopub.execute_input":"2024-12-02T13:11:58.197221Z","iopub.status.idle":"2024-12-02T13:12:35.601954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import lightning.pytorch as pl\nfrom typing import Tuple, List, Dict, Union\nfrom monai.networks.nets import UNet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric\nimport numpy as np\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)\n\n\nchannels = (48, 64, 80, 80)\nstrides_pattern = (2, 2, 1)       \nnum_res_units = 1\nlearning_rate = 1e-4\nnum_epochs = 50\n\nmodel = Model(spatial_dims=3, in_channels=1, out_channels=7, channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:12:35.605915Z","iopub.execute_input":"2024-12-02T13:12:35.606175Z","iopub.status.idle":"2024-12-02T13:13:03.163870Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nmodel = torch.load(\"/kaggle/input/3dunet-training/model.pth\")\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:03.165298Z","iopub.execute_input":"2024-12-02T13:13:03.166768Z","iopub.status.idle":"2024-12-02T13:13:04.656413Z"}},"outputs":[],"execution_count":null},{"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    print(total_overlap)\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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:04.658005Z","iopub.execute_input":"2024-12-02T13:13:04.658378Z","iopub.status.idle":"2024-12-02T13:13:04.670432Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:04.671381Z","iopub.execute_input":"2024-12-02T13:13:04.671669Z","iopub.status.idle":"2024-12-02T13:13:04.688769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Iterate over all train experiments and normalize using percentile 0.15 and 98. Then crop values from 0 to 1 and feed them into our model.\ntrain_data = []\ntrain_experiment_names = ['TS_5_4', 'TS_69_2', 'TS_6_6','TS_73_6', 'TS_86_3',\"TS_99_9\" ]\nvalid_experiment_names = ['TS_6_4']\nTRAIN_DIR = \"/kaggle/input/dataset-cryoet/\"\n\nfor name in train_experiment_names:\n    image = np.load(TRAIN_DIR + f\"train_image_{name}.npy\")\n    label = np.load(TRAIN_DIR + f\"train_label_{name}.npy\")\n    mean = image.flatten().mean()\n    std = image.flatten().std()\n    print(mean, std)\n    min_value = np.percentile(image.flatten(), 0.15)\n    max_value = np.percentile(image.flatten(), 98)\n    image = (image - min_value) / (max_value - min_value)    \n    image = np.clip(image, 0, 1)\n    train_data.append({\"image\": image, \"label\": label})\n    \nval_data = []\nfor name in valid_experiment_names:\n    image = np.load(TRAIN_DIR + f\"train_image_{name}.npy\")\n    label = np.load(TRAIN_DIR + f\"train_label_{name}.npy\")\n    min_value = np.percentile(image.flatten(), 0.15)\n    max_value = np.percentile(image.flatten(), 98)\n    image = (image - min_value) / (max_value - min_value)\n    image = (image - min_value) / (max_value - min_value)    \n    image = np.clip(image, 0, 1)\n    val_data.append({\"image\": image, \"label\": label})\n    print(np.unique(label))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:04.691316Z","iopub.execute_input":"2024-12-02T13:13:04.691600Z","iopub.status.idle":"2024-12-02T13:13:24.436678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image1 = train_data[0][\"image\"]\nlabel1 = train_data[0][\"label\"]\nprint(image1.shape,label1.shape)\n\npatches_image, coordinates_image  = extract_3d_patches_minimal_overlap([image1],96)\nlabel_image, coordinates_label = extract_3d_patches_minimal_overlap([label1],96)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:24.437575Z","iopub.execute_input":"2024-12-02T13:13:24.437874Z","iopub.status.idle":"2024-12-02T13:13:24.444128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(len(patches_image)):\n    patches_image[i] = np.expand_dims(patches_image[i], axis=0)\n\nprint(patches_image[0].shape)\nbatch_of_images = np.array(patches_image)\nprint(batch_of_images.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:24.445139Z","iopub.execute_input":"2024-12-02T13:13:24.445383Z","iopub.status.idle":"2024-12-02T13:13:24.594808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport gc\n\n# Check device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Move model to device and set to evaluation mode\nmodel = model.to(device)\nmodel.eval()\n\n# Batch size\nbatch_size = 10  # Adjust based on GPU memory\n\n# Store predictions\nall_predictions = []\n\n# Disable gradients for inference\nwith torch.no_grad():\n    for i in range(0, len(batch_of_images), batch_size):\n        # Slice the batch\n        batch = batch_of_images[i:i + batch_size]\n\n        # Convert batch to tensor and move to GPU\n        input_tensor = torch.from_numpy(batch).float().to(device)\n\n        # Make predictions\n        predictions = model(input_tensor)\n\n        # Store predictions (move them to CPU to save GPU memory)\n        all_predictions.append(predictions.cpu())\n\n        # Clear GPU cache\n        torch.cuda.empty_cache()\n        gc.collect()\n\n# Combine all predictions into a single tensor if needed\nfinal_predictions = torch.cat(all_predictions)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:24.596061Z","iopub.execute_input":"2024-12-02T13:13:24.596475Z","iopub.status.idle":"2024-12-02T13:13:33.671075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nprobabilities = F.softmax(final_predictions, dim=1)\npredicted_classes = torch.argmax(probabilities, dim=1)\nprint(predicted_classes.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:33.672150Z","iopub.execute_input":"2024-12-02T13:13:33.672421Z","iopub.status.idle":"2024-12-02T13:13:49.874893Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nfig, axes = plt.subplots(4,4,figsize=(15,15))\naxes = axes.flatten()\nfor i in range (len(axes)//2):\n    axes[i*2].imshow(patches_image[i][0,60,:,:])\n    axes[i*2+1].imshow(predicted_classes[i][60,:,:])\n    axes[i*2].set_title(f\"image {i}\")\n    axes[i*2+1].set_title(f\"label {i}\")\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:49.876073Z","iopub.execute_input":"2024-12-02T13:13:49.876327Z","iopub.status.idle":"2024-12-02T13:13:52.530810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfig,axes = plt.subplots(4,4,figsize=(15,15))\naxes = axes.flatten()\nfor i in range (len(axes)//2):\n    axes[i*2].imshow(patches_image[i][0,60,:,:])\n    axes[i*2+1].imshow(label_image[i][60,:,:])\n    axes[i*2].set_title(f\"image {i}\")\n    axes[i*2+1].set_title(f\"label {i}\")    \n        \nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:52.531782Z","iopub.execute_input":"2024-12-02T13:13:52.532034Z","iopub.status.idle":"2024-12-02T13:13:54.939029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BLOB_THRESHOLD = 500\nCERTAINTY_THRESHOLD = 0.5\n\nimport cc3d\nimport copick\n\nconfig_file = \"/kaggle/input/baseline-unet-train-submit/copick_test.config\"\nroot = copick.from_file(config_file)\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"\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\"}\nclasses = [1, 2, 3, 4, 5, 6]\nwith torch.no_grad():\n    location_df = []\n    for run in root.runs:\n        print(run)\n        tomo = run.get_voxel_spacing(10)\n        tomo = tomo.get_tomogram(tomo_type).numpy()\n        tomo_patches, coordinates  = extract_3d_patches_minimal_overlap([tomo], 96)\n        tomo_patched_data = [{\"image\": img} for img in tomo_patches]\n        tomo_ds = tomo_patched_data\n        #tomo_ds = CacheDataset(data=tomo_patched_data, transform=inference_transforms, cache_rate=1.0)\n        pred_masks = []\n        for i in range(len(tomo_ds)):\n            input_image = tomo_ds[i][\"image\"]\n            min_percentile = np.percentile(input_image.flatten(),0.15)\n            max_percentile = np.percentile(input_image.flatten(),98)\n            input_image = (input_image-min_percentile)/(max_percentile-min_percentile)\n            input_image = np.clip(input_image, 0, 1)\n            input_image = torch.from_numpy(input_image)\n            input_image = input_image.unsqueeze(0).unsqueeze(0)\n            input_tensor = input_image.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        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    \n    location_df = pd.concat(location_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:13:54.940163Z","iopub.execute_input":"2024-12-02T13:13:54.940478Z","iopub.status.idle":"2024-12-02T13:14:27.654696Z"}},"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,"execution":{"iopub.status.busy":"2024-12-02T13:14:27.655748Z","iopub.execute_input":"2024-12-02T13:14:27.656398Z","iopub.status.idle":"2024-12-02T13:14:27.668054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import cc3d\n# from scipy.ndimage import binary_erosion, generate_binary_structure, iterate_structure\n\n# def create_3d_disc(radius=2):\n#     # Start with a basic 3D connectivity structure\n#     # connectivity=1 gives 6-connectivity (face neighbors)\n#     # connectivity=2 gives 18-connectivity (face and edge neighbors)\n#     # connectivity=3 gives 26-connectivity (face, edge, and vertex neighbors)\n#     basic_structure = generate_binary_structure(3, 2)\n    \n#     # Iterate the structure to create a disc shape\n#     # iterations=radius-1 because generate_binary_structure creates a 3x3x3 structure\n#     disc = iterate_structure(basic_structure, radius-1)\n    \n#     return disc\n    \n# def run_submit(experiments):\n#     df = []\n#     experiment, particle_type, x, y, z = [], [] ,[] ,[] ,[]\n    \n#     for i, id in enumerate(experiments):\n#         print(i, id, '---------------')\n#         volume = read_one_data(id, static_dir=f'/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns')\n#         print(volume.shape)\n        \n#         # Extract patches with minimal overlap\n#         patches_volume, coordinates = extract_3d_patches_minimal_overlap([volume], 96)\n#         D, H, W = volume.shape\n        \n#         # Initialize arrays for accumulating predictions\n#         probability = np.zeros((7, D, H, W), dtype=np.float32)\n#         count = np.zeros((7, D, H, W), dtype=np.float32)\n#         n_per_batch = 5\n#         count_coordinates = 0\n        \n#         for j in range(0, len(patches_volume), n_per_batch):\n#             with torch.no_grad():\n#                 input_patches = patches_volume[j:j+n_per_batch]\n#                 input_patches = np.expand_dims(input_patches, axis=1)\n                \n#                 # Create both original and rotated versions of the patches\n#                 input_orig = torch.from_numpy(input_patches).float()\n#                 input_rot = torch.rot90(input_orig, k=1, dims=(3, 4))# Rotate in the spatial dimensions\n                \n#                 # Stack them in the batch dimension\n#                 input_combined = torch.cat([input_orig, input_rot], dim=0).to(device)\n\n                \n#                 # Get predictions for both original and rotated patches\n#                 predictions_combined = model(input_combined)\n#                 predictions_combined = torch.nn.functional.softmax(predictions_combined, dim=1)\n\n                \n#                 # Split predictions back into original and rotated\n#                 n_patches = len(input_patches)\n#                 pred_orig = predictions_combined[:n_patches]\n#                 pred_rot = predictions_combined[n_patches:]\n                \n#                 # Rotate back the rotated predictions\n#                 pred_rot = torch.rot90(pred_rot, k=-1, dims=(3, 4))\n                \n#                 # Average the predictions\n#                 predictions = (pred_orig + pred_rot) / 2\n#                 predictions = predictions.cpu()\n                \n#                 # Clear GPU memory\n#                 torch.cuda.empty_cache()\n#                 gc.collect()\n                \n#                 # Accumulate weighted predictions\n#                 for prediction in predictions:\n#                     left_point = coordinates[count_coordinates]                    \n#                     # Add weighted predictions to the appropriate location\n#                     probability[:, \n#                               left_point[0]:96+left_point[0],\n#                               left_point[1]:left_point[1]+96,\n#                               left_point[2]:left_point[2]+96] = prediction * weight\n                    \n#                     count[:, \n#                           left_point[0]:96+left_point[0],\n#                           left_point[1]:left_point[1]+96,\n#                           left_point[2]:left_point[2]+96] += weight\n                    \n#                     count_coordinates += 1\n                    \n#         # Average the probabilities using the accumulated weights\n#         probability_maps = probability / (count + 0.0001)\n#         structure = create_3d_disc(radius=1)  # Adjust radius as needed\n#         particles = [1,2,3,4,5,6]\n#         names = ['apo-ferritin', \"beta-amylase\", 'beta-galactosidase','ribosome','thyroglobulin','virus-like-particle']\n#         radius = [6,6.5,9,15,13,13.5]\n#         for idx in particles:\n        \n#             # Get probability map for this particle type and apply threshold\n#             prob_map = probability_maps[idx,:,:,:]\n#             binary_mask = prob_map > 0.05\n            \n#             binary_mask = binary_erosion(binary_mask, structure=structure, iterations=2)\n\n#             # Find connected components\n#             labels = cc3d.connected_components(binary_mask)\n#             stats = cc3d.statistics(labels)\n            \n#             # Get volumes of all components (excluding background)\n#             volumes = stats[\"voxel_counts\"][1:]  # Exclude background (label 0)\n            \n#             # Calculate volume threshold\n#             sphere_volume = (4 / 3) * np.pi * ((radius[idx-1] * 0.8) ** 3)\n            \n#             # Filter components based on volume threshold\n#             valid_indices = [i for i, v in enumerate(volumes) if v >= sphere_volume * 0.25]\n            \n#             if not valid_indices:\n#                 continue\n            \n#             # Get centroids for valid components\n#             centroids = stats['centroids'][1:]  # Skip background (label 0)\n#             valid_centroids = centroids[valid_indices]  # Select only valid centroids\n            \n#             if len(valid_centroids) > 0:\n#                 # Scale coordinates by voxel size (10nm in your case)\n#                 scaled_centroids = 10 * valid_centroids  # Adjust for voxel size\n                \n#                 # Add to lists\n#                 n_particles = len(valid_centroids)\n#                 experiment.extend([id] * n_particles)\n#                 particle_type.extend([names[idx-1]] * n_particles)\n#                 x.extend(scaled_centroids[:, 2])  # ZYX to XYZ conversion\n#                 y.extend(scaled_centroids[:, 1])\n#                 z.extend(scaled_centroids[:, 0])\n            \n#     return experiment, particle_type, x, y, z\n\n# experiment, particle_type, x, y, z = run_submit(SAMPLES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:14:27.669217Z","iopub.execute_input":"2024-12-02T13:14:27.669535Z","iopub.status.idle":"2024-12-02T13:14:27.679193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pd.DataFrame({\n#     'id':np.arange(len(experiment)),\n#     'experiment':experiment,\n#     'particle_type':particle_type,\n#     'x':x,\n#     'y':y,\n#     'z':z\n# }).to_csv('submission.csv',index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:14:27.680213Z","iopub.execute_input":"2024-12-02T13:14:27.680589Z","iopub.status.idle":"2024-12-02T13:14:27.694022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# path_to_overlay = \"/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns\"\n# import json\n# names = ['apo-ferritin', \"beta-amylase\", 'beta-galactosidase','ribosome','thyroglobulin','virus-like-particle']\n# experiments =[]\n# particle_names= []\n# x_particles=[]\n# y_particles=[]\n# z_particles=[]\n\n# for i,id in enumerate(SAMPLES):\n#     file = path_to_overlay+f\"/{id}/Picks/\"\n#     for particle_name in names:\n#         file_name = file + particle_name + \".json\"\n#         with open(file_name,\"r\") as f:\n#             data = json.load(f)\n#         data = data[\"points\"]\n#         x = [d[\"location\"][\"x\"] for d in data]\n#         y = [d[\"location\"][\"y\"] for d in data]\n#         z = [d[\"location\"][\"z\"] for d in data]\n#         experiments.extend([id]*len(x))\n#         particle_names.extend([particle_name]*len(x))\n#         x_particles.extend(x)\n#         y_particles.extend(y)\n#         z_particles.extend(z)\n\n    \n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:14:27.694903Z","iopub.execute_input":"2024-12-02T13:14:27.695147Z","iopub.status.idle":"2024-12-02T13:14:27.705856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pd.DataFrame({\n#     'id':np.arange(len(experiments)),\n#     'experiment':experiments,\n#     'particle_type':particle_names,\n#     'x':x_particles,\n#     'y':y_particles,\n#     'z':z_particles\n# }).to_csv('check_sub.csv',index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T13:14:27.706726Z","iopub.execute_input":"2024-12-02T13:14:27.706937Z","iopub.status.idle":"2024-12-02T13:14:27.722777Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ","metadata":{}},{"cell_type":"code","source":"","metadata":{"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}]}