{"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":10011619,"sourceType":"datasetVersion","datasetId":6144035},{"sourceId":10040238,"sourceType":"datasetVersion","datasetId":6184828},{"sourceId":206165222,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"mode=\"local\"\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 optuna-integration[pytorch-lightning]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:13:58.632082Z","iopub.execute_input":"2024-12-04T07:13:58.632408Z","iopub.status.idle":"2024-12-04T07:15:00.142506Z","shell.execute_reply.started":"2024-12-04T07:13:58.632377Z","shell.execute_reply":"2024-12-04T07:15:00.141639Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:15:52.296768Z","iopub.execute_input":"2024-12-04T07:15:52.297149Z","iopub.status.idle":"2024-12-04T07:16:28.083646Z","shell.execute_reply.started":"2024-12-04T07:15:52.297118Z","shell.execute_reply":"2024-12-04T07:16:28.082463Z"}},"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    # 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,"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:16:28.085623Z","iopub.execute_input":"2024-12-04T07:16:28.086758Z","iopub.status.idle":"2024-12-04T07:16:28.106285Z","shell.execute_reply.started":"2024-12-04T07:16:28.086724Z","shell.execute_reply":"2024-12-04T07:16:28.105376Z"}},"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":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:16:28.107787Z","iopub.execute_input":"2024-12-04T07:16:28.108194Z","iopub.status.idle":"2024-12-04T07:16:28.131218Z","shell.execute_reply.started":"2024-12-04T07:16:28.108154Z","shell.execute_reply":"2024-12-04T07:16:28.130311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nTRAIN_DIR = \"/kaggle/input/dataset-cryoet/\"\nTEST_DATA_DIR = \"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:16:28.133257Z","iopub.execute_input":"2024-12-04T07:16:28.133606Z","iopub.status.idle":"2024-12-04T07:16:28.145754Z","shell.execute_reply.started":"2024-12-04T07:16:28.133562Z","shell.execute_reply":"2024-12-04T07:16:28.144974Z"}},"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.\nimport numpy as np\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']\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(), 99)\n    #image = (image - min_value) / (max_value - min_value)    \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(), 99)\n    #image = (image - min_value) / (max_value - min_value)\n    val_data.append({\"image\": image, \"label\": label})\n    print(np.unique(label))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-04T07:16:49.954166Z","iopub.execute_input":"2024-12-04T07:16:49.955088Z","iopub.status.idle":"2024-12-04T07:17:08.221978Z","shell.execute_reply.started":"2024-12-04T07:16:49.955033Z","shell.execute_reply":"2024-12-04T07:17:08.221062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport scipy.ndimage\nfrom typing import Dict, List, Tuple, Optional\nfrom multiprocessing import Pool, cpu_count\nfrom functools import partial\n\nclass ParticleExtractor:\n    def __init__(self, margin: int = 0, n_classes: int = 7):\n        self.margin = margin\n        self.n_classes = n_classes\n        self.class_radius = {1:6, 2:6.5, 3:9, 4:15, 5:13, 6:13.5}\n        self.particle_labels = {i:[] for i in range(1,n_classes)}\n        self.particle_images = {i:[] for i in range(1,n_classes)}\n        #we will create a list of dicts with the counts for each class\n        self.particle_counts = []    \n        \n    @staticmethod\n    def _extract_class_particles(args):\n        \"\"\"\n        Helper method to extract particles for a single class.\n        \"\"\"\n        class_id, valid_region, valid_image, radius, particles_per_class, crop_margin = args\n        \n        particles_found = 0\n        images = []\n        masks = []\n        \n        # Label connected components for this class\n        labeled_mask, num_features = scipy.ndimage.label(valid_region == class_id)\n        \n        # Shuffle particle indices to randomize selection\n        particle_indices = np.random.permutation(range(1, num_features + 1))\n        \n        for particle_id in particle_indices:\n            if particles_found >= particles_per_class:\n                break\n                \n            where_particle = np.where(labeled_mask == particle_id)\n            if len(where_particle[0]) == 0:\n                continue\n            \n            # Get particle bounds\n            z, y, x = where_particle\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            length_z = max_z - min_z\n            length_y = max_y - min_y\n            length_x = max_x - min_x\n            \n            radius_masked = 0.5 * min(length_z, length_x, length_y)\n            \n            if radius-radius_masked > 2:\n                continue\n            \n            # Add margin around particle\n            min_z = max(0, min_z - crop_margin)\n            max_z = min(labeled_mask.shape[0], max_z + crop_margin + 1)\n            min_y = max(0, min_y - crop_margin)\n            max_y = min(labeled_mask.shape[1], max_y + crop_margin + 1)\n            min_x = max(0, min_x - crop_margin)\n            max_x = min(labeled_mask.shape[2], max_x + crop_margin + 1)\n            \n            # Extract crop region\n            particle_region = np.zeros_like(labeled_mask)\n            particle_region[min_z:max_z, min_y:max_y, min_x:max_x] = (\n                labeled_mask[min_z:max_z, min_y:max_y, min_x:max_x] == particle_id\n            )\n            \n            # Get final crop\n            particle_mask = particle_region[min_z:max_z, min_y:max_y, min_x:max_x] * class_id\n            particle_img = valid_image[min_z:max_z, min_y:max_y, min_x:max_x]\n            \n            masks.append(particle_mask)\n            images.append(particle_img)\n            particles_found += 1\n            \n        return class_id, masks, images\n\n    def extract_particles(self, data: List[Dict[str, np.ndarray]], \n                         particles_per_class: int) -> None:\n        for i, dct in enumerate(data):\n            \n            self.particle_counts.append({i:0 for i in range(1,self.n_classes)})\n            \n            image, mask = dct[\"image\"], dct[\"label\"]\n                        \n            # Remove edge regions\n            z_max, y_max, x_max = mask.shape\n            valid_region = mask[\n            self.margin:z_max-self.margin,\n            self.margin:y_max-self.margin,\n            self.margin:x_max-self.margin\n            ]\n            valid_image = image[\n                self.margin:z_max-self.margin,\n                self.margin:y_max-self.margin,\n                self.margin:x_max-self.margin\n            ]\n            \n            # Prepare arguments for parallel processing\n            class_ids = np.unique(valid_region)[1:]  # Skip background\n            args_list = [\n                (class_id, \n                valid_region, \n                valid_image, \n                self.class_radius[class_id]*0.8,\n                particles_per_class,\n                2)  # crop_margin\n                for class_id in class_ids\n            ]\n            \n            # Process classes in parallel\n            n_cores = cpu_count()\n            with Pool(processes=n_cores) as pool:\n                results = pool.map(self._extract_class_particles, args_list)\n            \n            # Store results\n            for class_id, masks, images in results:\n                self.particle_labels[class_id].extend(masks)\n                self.particle_images[class_id].extend(images)\n                self.particle_counts[i][class_id] += len(masks)\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:26.665030Z","iopub.execute_input":"2024-12-03T16:42:26.665750Z","iopub.status.idle":"2024-12-03T16:42:26.681219Z","shell.execute_reply.started":"2024-12-03T16:42:26.665702Z","shell.execute_reply":"2024-12-03T16:42:26.680196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\nwith open(\"/kaggle/input/cache-particles/cache_particles.pkl\",\"rb\") as file:\n    cache_particles = pickle.load(file)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:26.682530Z","iopub.execute_input":"2024-12-03T16:42:26.683277Z","iopub.status.idle":"2024-12-03T16:42:26.810672Z","shell.execute_reply.started":"2024-12-03T16:42:26.683224Z","shell.execute_reply":"2024-12-03T16:42:26.809602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cache_particles.particle_counts[2]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:26.811996Z","iopub.execute_input":"2024-12-03T16:42:26.812319Z","iopub.status.idle":"2024-12-03T16:42:26.819186Z","shell.execute_reply.started":"2024-12-03T16:42:26.812281Z","shell.execute_reply":"2024-12-03T16:42:26.818327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n# Calculate the total number of rows needed (one per class)\nn_rows = 6  # for classes 1-6\nn_cols = 6  # for your image pairs\n\n# Create one large figure\nplt.figure(figsize=(20, 15))  # Adjust size as needed\n\n# Loop through each class and image\nfor class_idx in range(1, 7):  # classes 1-6\n    for j in range(n_cols):\n        # Calculate subplot position\n        plot_idx = (class_idx-1) * n_cols + j + 1\n        \n        plt.subplot(n_rows, n_cols, plot_idx)\n        \n        if j % 2:\n            im = plt.imshow(cache_particles.particle_images[class_idx][j][5,:,:])\n            plt.colorbar(im, shrink=0.5)\n            plt.title(f\"Image {j}\")\n        else:\n            plt.imshow(cache_particles.particle_labels[class_idx][j][5,:,:])\n            plt.title(f\"Mask {j}\")\n            \n        # Add titles/labels as needed\n        plt.axis('off')  # Optional: remove axes\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:26.820367Z","iopub.execute_input":"2024-12-03T16:42:26.820749Z","iopub.status.idle":"2024-12-03T16:42:30.411073Z","shell.execute_reply.started":"2024-12-03T16:42:26.820709Z","shell.execute_reply":"2024-12-03T16:42:30.410131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.transforms import Transform\nfrom typing import List, Dict\nimport numpy as np\nimport torch\nfrom collections import Counter\n\nclass PasteParticlesd(Transform):\n    def __init__(self, keys: List[str], label_key: str, cache_particles):\n        # Initialize transform parameters\n        self.keys = keys\n        self.label_key = label_key\n        \n        # Convert particle cache to CPU numpy arrays immediately\n        self.particle_labels = {\n            k: [np.asarray(p.cpu().numpy() if torch.is_tensor(p) else p) \n                for p in v]\n            for k, v in cache_particles.particle_labels.items()\n        }\n        self.particle_images = {\n            k: [np.asarray(p.cpu().numpy() if torch.is_tensor(p) else p) \n                for p in v]\n            for k, v in cache_particles.particle_images.items()\n        }\n\n    def __call__(self, data: Dict) -> Dict:\n        # Skip transform with 50% probability\n        if np.random.rand() < 0.5:\n            return data\n            \n        # Get image and label data, ensuring they're on CPU as numpy arrays\n        label = data[self.label_key]\n        image = data[self.keys[0]]\n        \n                # Convert to numpy arrays\n        label_np = label.cpu().numpy() if torch.is_tensor(label) else label\n        image_np = image.cpu().numpy() if torch.is_tensor(image) else image\n        \n        # Remove batch dimension if present\n        label_np = np.squeeze(label_np, axis=0)\n        image_np = np.squeeze(image_np, axis=0)\n        \n        # Get dimensions for particle placement\n        max_z, max_y, max_x = label_np.shape\n        \n        # Randomly select particle classes to process\n        indices = np.random.choice(range(6), 6, replace=False)\n        \n        # Process each selected particle class\n        for class_idx in indices:\n            # Skip with 50% probability\n            if np.random.rand() < 0.5:\n                continue\n                \n            # Get available particles for this class\n            class_particles = self.particle_labels[class_idx + 1]\n            if not class_particles:  # Skip if no particles available\n                continue\n                \n            # Select random particle\n            particle_idx = np.random.choice(len(class_particles))\n            particle_label = class_particles[particle_idx]\n            particle_image = self.particle_images[class_idx + 1][particle_idx]\n            \n            for _ in range(50):\n                # Calculate valid placement range\n                if max_z-particle_label.shape[0]<0 or max_y-particle_label.shape[1]<0 or max_x - particle_label.shape[2]<0:\n                    continue\n                \n                z_start = np.random.randint(0, np.floor(max_z - particle_label.shape[0]))\n                y_start = np.random.randint(0, np.floor(max_y - particle_label.shape[1]))\n                x_start = np.random.randint(0, np.floor(max_x - particle_label.shape[2]))\n\n                \n                # Check if space is empty\n                block_try = label_np[\n                    z_start:z_start + particle_label.shape[0],\n                    y_start:y_start + particle_label.shape[1],\n                    x_start:x_start + particle_label.shape[2]\n                ]\n                \n                if np.sum(block_try) == 0:\n                    \n                    # Place particle\n                    label_np[\n                        z_start:z_start + particle_label.shape[0],\n                        y_start:y_start + particle_label.shape[1],\n                        x_start:x_start + particle_label.shape[2]\n                    ] = particle_label\n                    \n                    image_np[\n                        z_start:z_start + particle_image.shape[0],\n                        y_start:y_start + particle_image.shape[1],\n                        x_start:x_start + particle_image.shape[2]\n                    ] = particle_image\n                    \n                    break\n        \n        # Restore batch dimension\n        label_np = np.expand_dims(label_np, axis=0)\n        image_np = np.expand_dims(image_np, axis=0)\n        \n        # Convert back to original format (tensor or numpy)\n        data[self.label_key] = torch.from_numpy(label_np).to(device=label.device if torch.is_tensor(label) else 'cpu')\n        data[self.keys[0]] = torch.from_numpy(image_np).to(device=image.device if torch.is_tensor(image) else 'cpu') \n                   \n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:30.413662Z","iopub.execute_input":"2024-12-03T16:42:30.413941Z","iopub.status.idle":"2024-12-03T16:42:30.427666Z","shell.execute_reply.started":"2024-12-03T16:42:30.413914Z","shell.execute_reply":"2024-12-03T16:42:30.426683Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import lightning.pytorch as pl\nfrom typing import Union\nfrom monai.networks.nets import UNet,AttentionUnet\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        optimizer = \"Adam\"):\n    \n        super().__init__()\n        #saves hyperparameters\n        self.save_hyperparameters()\n\n        #sets model UNet based on hyperparameters\n        self.model = AttentionUnet(\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            dropout=0.2\n        )\n\n        #loss function\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, ignore_empty=True, reduction=None)\n        \n        self.train_loss = 0\n        self.val_metric = []\n        self.num_train_batch = 0\n        self.num_val_batch = 0\n        self.val_loss = 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        print(x.shape)\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():\n            x, y = batch['image'], batch['label']\n            y_hat = self(x)\n            val_loss_value = self.loss_fn(y_hat, y)\n            self.val_loss += val_loss_value\n            #Does argmax (reduces dimensions) to get the max class at each pixel and then one-hots (one-hots back again)\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            self.metric_fn(y_pred=metric_val_outputs, y=metric_val_labels)\n            metrics = self.metric_fn.get_buffer()\n            #metrics = torch.nan_to_num(metrics, nan=0.0)\n            self.val_metric.append(metrics) \n            self.num_val_batch += 1\n            return {'val_metrics': metrics}\n    \n    def on_validation_epoch_end(self):\n        all_metrics = torch.cat(self.val_metric, dim=0)  \n        mean_metrics = torch.nanmean(all_metrics, dim=0)\n        print(f\"Mean Dice per class: {mean_metrics}\")\n        mean_dice = torch.mean(mean_metrics)\n        val_loss_average = self.val_loss/self.num_val_batch\n        print(f\"Loss: {val_loss_average}\")\n        self.log(\"val_loss\",val_loss_average, prog_bar=True)\n        self.log(\"val_dice\", mean_dice)\n        self.val_metric = []\n        self.val_loss = 0\n        self.num_val_batch = 0\n        \n    def configure_optimizers(self):\n        optimizers = {\n            'Adam': torch.optim.Adam,\n            'AdamW': torch.optim.AdamW,\n            'SGD': torch.optim.SGD\n        }\n        opt = optimizers[self.hparams.optimizer](self.parameters(), lr=self.hparams.lr)\n        return opt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T16:42:30.428808Z","iopub.execute_input":"2024-12-03T16:42:30.429088Z","iopub.status.idle":"2024-12-03T16:42:30.705125Z","shell.execute_reply.started":"2024-12-03T16:42:30.429061Z","shell.execute_reply":"2024-12-03T16:42:30.704376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from lightning.pytorch.callbacks import EarlyStopping\nfrom optuna.integration.pytorch_lightning import PyTorchLightningPruningCallback\nfrom monai.transforms import RandZoomd\n\ndef set_training(n_samples,dim):\n    \n    non_random_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n        NormalizeIntensityd(keys=\"image\")\n    ])\n\n    # Create the cached dataset for non-random transforms\n    # This stores the transformed data in CPU memory\n    raw_train_ds = CacheDataset(\n        data=train_data,\n        transform=non_random_transforms,\n        cache_rate=1.0  # Cache all data in memory\n    )\n    \n    # Now set up the random transforms that are applied during training\n    # Note that we pass the particle_extractor directly to PasteParticlesd\n    random_transforms = Compose([\n        RandCropByLabelClassesd(\n            keys=[\"image\", \"label\"],\n            label_key=\"label\",\n            spatial_size=[dim, dim, dim],\n            num_classes=7,\n            num_samples=n_samples\n        ),\n        PasteParticlesd(\n            keys=[\"image\", \"label\"],\n            label_key=\"label\",\n            cache_particles=cache_particles  # Pass the original cache_particles\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        RandZoomd(keys=[\"image\", \"label\"], prob=0.5, min_zoom=1, max_zoom=1.2)\n    ])\n    \n    # Create the final dataset with random transforms\n    train_ds = Dataset(\n        data=[d for d in raw_train_ds],\n        transform=random_transforms\n    )\n    \n    # Set up the DataLoader with optimized settings\n    train_dataloader = DataLoader(\n        train_ds,\n        batch_size=1,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=True\n    )\n\n    return train_dataloader\n\n\ndef set_validation(dim):\n\n    non_random_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n        NormalizeIntensityd(keys=\"image\")\n    ])\n\n    val_images,val_labels = [dcts['image'] for dcts in val_data],[dcts['label'] for dcts in val_data]\n    \n    val_image_patches, _ = extract_3d_patches_minimal_overlap(val_images, dim)\n    val_label_patches, _ = extract_3d_patches_minimal_overlap(val_labels, dim)\n    \n    val_patched_data = [{\"image\": img, \"label\": lbl} for img, lbl in zip(val_image_patches, val_label_patches)]\n    \n    valid_ds = CacheDataset(data=val_patched_data, transform=non_random_transforms, cache_rate=1.0)\n    \n    valid_batch_size = 16\n    # DataLoader remains the same\n    valid_loader = DataLoader(\n        valid_ds,\n        batch_size=valid_batch_size,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True\n    )\n\n    return valid_loader\n    \n\ndef objective(trial):\n\n    #define parameters suggested\n    #num_samples = trial.suggest_int('num_samples', 2, 25)\n    \n    lr = trial.suggest_float('lr', 1e-5, 1e-2, log=True)\n    #optimizer = trial.suggest_categorical('optimizer', ['Adam', 'AdamW', 'SGD'])\n    #weight_decay = trial.suggest_float('weight_decay', 1e-6, 1e-3, log=True)\n    \n    #fixed parameters\n    spatial_dims = 3\n    dim = 64\n    in_channels, out_channels = 1, 7\n    channels = (48, 64, 80, 80)\n    strides_pattern = (2, 2, 1)       \n    num_res_units = 1\n    scheduler = \"CosineAnnealingLR\"\n    optimizer = \"Adam\"\n    num_samples = 14\n    \n    #get dataloaders\n    train_dataloader = set_training(num_samples,dim)\n    val_dataloader = set_validation(dim)\n\n    model= Model(spatial_dims, in_channels, out_channels, channels, strides_pattern, num_res_units, lr, optimizer)\n\n    trainer = pl.Trainer(\n        max_epochs=15,\n        accelerator=\"gpu\",\n        devices=1,\n        callbacks=[\n            EarlyStopping(monitor=\"val_loss\", mode=\"min\", patience=2),\n            PyTorchLightningPruningCallback(trial, monitor=\"val_loss\")\n        ]\n    )\n    \n    trainer.fit(model, train_dataloader, val_dataloader)\n    return trainer.callback_metrics[\"val_loss\"].item()\n\n\nimport optuna\n\n# Create study that maximizes validation dice score\nstudy = optuna.create_study(direction=\"minimize\")\nn_trials = 50  # Number of parameter combinations to try\n\n# Run optimization\nstudy.optimize(objective, n_trials=n_trials)\n\nprint(\"Best trial:\")\ntrial = study.best_trial\nprint(f\"  Value: {trial.value}\")\nprint(\"  Params: \")\nfor key, value in trial.params.items():\n    print(f\"    {key}: {value}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T17:07:52.770942Z","iopub.execute_input":"2024-12-03T17:07:52.771347Z","iopub.status.idle":"2024-12-03T17:09:22.569144Z","shell.execute_reply.started":"2024-12-03T17:07:52.771310Z","shell.execute_reply":"2024-12-03T17:09:22.564862Z"}},"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(spatial_dims=3, in_channels=1, out_channels=7, channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)","metadata":{"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')\nfrom lightning.pytorch.callbacks import TQDMProgressBar\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    accelerator='gpu',\n    devices=1,\n    log_every_n_steps=1,  \n    enable_progress_bar=True,\n    callbacks=[TQDMProgressBar(refresh_rate=1)]\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ntrainer.fit(model, train_dataloader, valid_loader)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.save_checkpoint(\"model500epochs.ckpt\")\ntorch.save(model,\"model.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}