{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":31420.258385,"end_time":"2025-02-02T03:55:57.566510","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-02-01T19:12:17.308125","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"25433451","cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport json\nimport torch\nimport torch.nn as nn\nimport gc\nimport random\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom scipy.optimize import linear_sum_assignment\nfrom IPython.display import display, clear_output\nfrom torch.optim.lr_scheduler import ExponentialLR\ntry :\n    import zarr\n    import monai\n    import cc3d\n    import torch_optimizer as optim\n    from monai.networks.blocks import MaxAvgPool\nexcept :\n    !pip install zarr\n    !pip install monai\n    !pip install segmentation_models_pytorch\n    !pip install --no-index --find-links=/kaggle/input/hengck-czii-cryo-et-01/wheel_file connected-components-3d\n    !pip install torch-optimizer\n    import zarr\n    import monai\n    import cc3d\n    import torch_optimizer as optim\n    from monai.networks.blocks import MaxAvgPool","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2025-02-01T19:12:19.848549Z","iopub.status.busy":"2025-02-01T19:12:19.848230Z","iopub.status.idle":"2025-02-01T19:13:13.499535Z","shell.execute_reply":"2025-02-01T19:13:13.498745Z"},"papermill":{"duration":53.658982,"end_time":"2025-02-01T19:13:13.501086","exception":false,"start_time":"2025-02-01T19:12:19.842104","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5182c414","cell_type":"code","source":"DATA_KAGGLE_DIR = '/kaggle/input/czii-cryo-et-object-identification'\nTRAIN_DIR = f'{DATA_KAGGLE_DIR}/train'\nTEST_DIR = f'{DATA_KAGGLE_DIR}/test'\n\nTRAIN_EXP = [\"TS_5_4\",\"TS_69_2\",\"TS_6_4\",\"TS_6_6\",\"TS_73_6\",\"TS_86_3\",\"TS_99_9\"]\nTEST_EXP = [\"TS_5_4\",\"TS_69_2\",\"TS_6_4\"]\n\nscale = 10.012444196428572\n\nPATCH_SIZE = 64\n\nOBJECT_DICT = {\n    'apo-ferritin': {'label': 1, 'radius': 60/scale}, \n    'beta-galactosidase': {'label': 2, 'radius': 90/scale}, \n    'ribosome': {'label': 3, 'radius': 150/scale}, \n    'thyroglobulin': {'label': 4, 'radius': 130/scale}, \n    'virus-like-particle': {'label': 5, 'radius': 135/scale},\n    #'beta-amylase' : {'label': 6, 'radius':65/scale},\n}\n\ndef read_one_data(id, static_dir):\n    zarr_dir = f'{static_dir}/{id}/VoxelSpacing10.000'\n    zarr_file = f'{zarr_dir}/denoised.zarr'\n    zarr_data = zarr.open(zarr_file, mode='r')\n    volume = zarr_data[0][:]\n    # mean = volume.mean()\n    # std = volume.std()\n    # volume = (volume - mean) / std\n    return volume\n\n\ndef read_one_truth(id, overlay_dir):\n    location={}\n    json_dir = f'{overlay_dir}/{id}/Picks'\n    for p in OBJECT_DICT.keys():\n        json_file = f'{json_dir}/{p}.json'\n        with open(json_file, 'r') as f:\n            json_data = json.load(f)\n\n        num_point = len(json_data['points'])\n        loc = [list(reversed(list(json_data['points'][i]['location'].values())))  for i in range(num_point)]\n        location[p] = [[coo/scale for coo in coos] for coos in loc ]\n    return location\n\n\n\ndef visualize(data):\n    \"\"\"\n    Visualize multiple images or masks side by side for each slice in the dataset.\n    \n    Parameters:\n    data (dict): Dictionary where keys are the names of the data items,\n                 and values are either 3D images or masks.\n    \"\"\"\n    # Ensure all data items have the same depth (number of slices)\n    z_sizes = {name: item.shape[0] for name, item in data.items() if name!=\"center\"}\n    if len(set(z_sizes.values())) != 1:\n        raise ValueError(\"All items must have the same number of slices along the Z-axis.\")\n    \n    z_size = next(iter(z_sizes.values()))  # Depth of slices (Z-axis)\n    keys = list(data.keys())  # Get all keys for consistent ordering\n    # Loop over each slice (Z-index)\n    for z in range(z_size):\n        fig, axes = plt.subplots(1, len(keys), figsize=(3 * len(keys), 5))\n        for idx, key in enumerate(keys):\n            item = data[key]\n            cmap = \"gray\" if \"volume\" in key.lower() else \"jet\"\n            \n            # Display the slice\n            axes[idx].imshow(item[z, :, :], cmap=cmap)\n            axes[idx].set_title(f\"{key} - Slice (Z={z})\")\n            axes[idx].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n\ndef draw_cylinder_in_image_fast(image, center, radius, z_factor, yx_factor, value):\n    new_radius = radius * yx_factor\n    half_height = round(radius * z_factor)\n    \n    z_min = max(round(center[0] - half_height), 0)\n    y_min = max(round(center[1] - new_radius), 0)\n    x_min = max(round(center[2] - new_radius), 0)\n\n    z_max = min(round(center[0] + half_height) + 1, image.shape[0])\n    y_max = min(round(center[1] + new_radius) + 1, image.shape[1])\n    x_max = min(round(center[2] + new_radius) + 1, image.shape[2])\n    \n    local_region = image[z_min:z_max, y_min:y_max, x_min:x_max]\n    local_center = (half_height, new_radius, new_radius)\n    \n    local_region = draw_cylinder_in_local_image(local_region, local_center, new_radius, half_height, value)\n    image[z_min:z_max, y_min:y_max, x_min:x_max] = np.bitwise_or(local_region,image[z_min:z_max, y_min:y_max, x_min:x_max])\n\n    return image\n\ndef draw_cylinder_in_local_image(image, center, radius, half_height, value):\n    shape = image.shape\n\n    z, y, x = np.indices(shape)\n\n    distance = (y - center[1])**2 + (x - center[2])**2\n\n    z_bounds = (z >= center[0] - half_height) & (z <= center[0] + half_height)\n    yx_bounds = distance <= radius**2\n    \n    cylinder = z_bounds & yx_bounds\n\n    image[cylinder] = value\n\n    return image\n\ndef draw_sphere_in_image_fast(image, center, radius, radius_factor, value):\n    new_radius = radius * radius_factor\n    z_min = max(round(center[0] - new_radius), 0)\n    y_min = max(round(center[1] - new_radius), 0)\n    x_min = max(round(center[2] - new_radius), 0)\n\n    z_max = min(round(center[0] + new_radius) + 1, image.shape[0])\n    y_max = min(round(center[1] + new_radius) + 1, image.shape[1])\n    x_max = min(round(center[2] + new_radius) + 1, image.shape[2])\n    \n    local_region = image[z_min:z_max, y_min:y_max, x_min:x_max]\n    local_center = (new_radius, new_radius, new_radius)\n    \n    local_region = draw_sphere_in_local_image(local_region, local_center, new_radius, value)\n    image[z_min:z_max, y_min:y_max, x_min:x_max] = np.bitwise_or(local_region,image[z_min:z_max, y_min:y_max, x_min:x_max])\n\n    return image\n\ndef draw_sphere_in_local_image(image, center, radius, value):\n    shape = image.shape\n\n    z, y, x = np.indices(shape)\n\n    distance = (z - center[0])**2 + (y - center[1])**2 + (x - center[2])**2\n    \n    cylinder = distance <= radius**2\n\n    image[cylinder] = value\n\n    return image\n\ndef draw_dentroid(image, center,value):\n\n    z, y, x = map(round, center)\n    \n    z_range = [z-1, z + 1]\n    y_range = [y-1, y + 1]\n    x_range = [x-1, x + 1]\n    \n    for zi in z_range:\n        for yi in y_range:\n            for xi in x_range:\n                # Ensure the indices are within the bounds of the image\n                if 0 <= zi < image.shape[0] and 0 <= yi < image.shape[1] and 0 <= xi < image.shape[2]:\n                    image[zi, yi, xi] = value\n\n    return image\n\ndef generate_heatmap(shape, centers, sigma=2):\n    heatmap = np.zeros(shape)  # Initialize the 3D heatmap\n\n    for center in centers:\n        z, y, x = center\n        radius = 3 * sigma  # Define a cutoff radius (3 sigma covers ~99.7% of Gaussian)\n        \n        # Define bounds for the local region\n        z_min = max(round(z - radius), 0)\n        y_min = max(round(y - radius), 0)\n        x_min = max(round(x - radius), 0)\n\n        z_max = min(round(z + radius) + 1, shape[0])\n        y_max = min(round(y + radius) + 1, shape[1])\n        x_max = min(round(x + radius) + 1, shape[2])\n\n        # Extract the local region\n        local_region = heatmap[z_min:z_max, y_min:y_max, x_min:x_max]\n        local_center = (z - z_min, y - y_min, x - x_min)  # Local center coordinates\n\n        # Generate Gaussian blob in the local region\n        heatmap[z_min:z_max, y_min:y_max, x_min:x_max] += generate_local_gaussian(\n            local_region.shape, local_center, sigma\n        )\n\n    return heatmap\n\n\ndef generate_local_gaussian(shape, center, sigma):\n    z, y, x = np.indices(shape)  # Create coordinate grids for the local region\n\n    # Compute the squared distance from the center\n    distance = (z - center[0])**2 + (y - center[1])**2 + (x - center[2])**2\n\n    # Compute the Gaussian values\n    gaussian = np.exp(-distance / (2 * sigma**2))\n\n    return gaussian\n\ndef generate_centers(image_shape, patch_size):\n    \n    centers = [\n        [z, y, x]\n        for z in range(patch_size[0]//2, image_shape[0], 3*patch_size[0] // 4)\n        for y in range(patch_size[1]//2, image_shape[1], 3*patch_size[1] // 4)\n        for x in range(patch_size[2]//2, image_shape[2], 3*patch_size[2] // 4) \n        if z < image_shape[0]-patch_size[0]//2 and y < image_shape[1]-patch_size[1]//2 and x < image_shape[2]-patch_size[2]//2\n    ]\n    centers.append([c-patch_size[i] for i,c in enumerate(image_shape)])\n    return centers\n    \ndef crop_with_center(image, patch_size, center):\n\n    center = [round (c) for c in center]\n    dims = image.shape\n    if len(dims) == 3:\n        start_indices = [max(0, c - p // 2) for c, p in zip(center, patch_size)]\n        end_indices = [min(dim, c + (p + 1) // 2) for c, p, dim in zip(center, patch_size, dims)]\n        \n        cropped_patch = image[\n            start_indices[0]:end_indices[0],\n            start_indices[1]:end_indices[1],\n            start_indices[2]:end_indices[2]\n        ]\n    else:\n        start_indices = [max(0, c - p // 2) for c, p in zip(center, patch_size)]\n        end_indices = [min(dim, c + (p + 1) // 2) for c, p, dim in zip(center, patch_size, dims[1:])]\n        \n        cropped_patch = image[\n            :,\n            start_indices[0]:end_indices[0],\n            start_indices[1]:end_indices[1],\n            start_indices[2]:end_indices[2]\n        ]\n    return cropped_patch\n    \ndef pad(result, patch_size):\n    shape = np.array(result[\"volume\"].shape)\n    expected_shape = np.array(patch_size)\n    diff = expected_shape - shape\n\n    if np.all(diff == 0):\n        return result\n    else:\n        # Compute mean and std of the original volume\n        mean = np.mean(result[\"volume\"])\n        std = np.std(result[\"volume\"])\n\n        # Create a fill matrix with the same mean and std\n        fill_matrix = np.random.normal(mean, std, size=tuple(expected_shape)).astype(result[\"volume\"].dtype)\n\n        # Determine random start points for padding\n        zyx_start = [np.random.randint(0, diff[i] + 1) for i in range(3)]\n\n        # Fill the padded matrix with the original volume\n        fill_matrix[\n            zyx_start[0]:zyx_start[0] + shape[0],\n            zyx_start[1]:zyx_start[1] + shape[1],\n            zyx_start[2]:zyx_start[2] + shape[2]\n        ] = result[\"volume\"]\n        result[\"volume\"] = fill_matrix\n\n        # Pad other keys in the result\n        for key in result.keys():\n            if key == \"volume\":\n                continue\n\n            current_shape = result[key].shape\n            if len(current_shape) > 3:\n                mask_padder = np.zeros((current_shape[0], *expected_shape), dtype=bool)\n                mask_padder[\n                    :,\n                    zyx_start[0]:zyx_start[0] + shape[0],\n                    zyx_start[1]:zyx_start[1] + shape[1],\n                    zyx_start[2]:zyx_start[2] + shape[2]\n                ] = result[key] if key != \"containers\" else result[key] + 0.5\n                result[key] = mask_padder\n\n            else:\n                mask_padder = np.zeros(expected_shape, dtype=bool)\n                mask_padder[\n                    zyx_start[0]:zyx_start[0] + shape[0],\n                    zyx_start[1]:zyx_start[1] + shape[1],\n                    zyx_start[2]:zyx_start[2] + shape[2]\n                ] = result[key] if key != \"containers\" else result[key] + 0.5\n                result[key] = mask_padder\n\n    return result\n\ndef pad_tensor(result, patch_size):\n    shape = torch.tensor(result[\"volume\"].shape)\n    expected_shape = torch.tensor(patch_size)\n    diff = expected_shape - shape\n\n    if (diff == 0).all():\n        return result\n    else:\n        # Compute mean and std of the original volume\n        mean = result[\"volume\"].mean()\n        std = result[\"volume\"].std()\n\n        # Create a fill matrix with the same mean and std\n        fill_matrix = torch.normal(mean, std, size=tuple(expected_shape), dtype=result[\"volume\"].dtype)\n\n        # Determine random start points for padding\n        zyx_start = [torch.randint(0, diff[i].item() + 1, (1,)).item() for i in range(3)]\n\n        # Fill the padded matrix with the original volume\n        fill_matrix[\n            zyx_start[0]:zyx_start[0] + shape[0],\n            zyx_start[1]:zyx_start[1] + shape[1],\n            zyx_start[2]:zyx_start[2] + shape[2]\n        ] = result[\"volume\"]\n        result[\"volume\"] = fill_matrix\n\n        # Pad other keys in the result\n        for key in result.keys():\n            if key == \"volume\":\n                continue\n\n            current_shape = result[key].shape\n            if len(current_shape) > 3:\n                mask_padder = torch.zeros(\n                    (current_shape[0], *expected_shape), dtype=torch.bool\n                )\n                mask_padder[\n                    :,\n                    zyx_start[0]:zyx_start[0] + shape[0],\n                    zyx_start[1]:zyx_start[1] + shape[1],\n                    zyx_start[2]:zyx_start[2] + shape[2]\n                ] = result[key] if key != \"containers\" else result[key] + 0.5\n                result[key] = mask_padder\n\n            else:\n                mask_padder = torch.zeros(list(expected_shape), dtype=torch.bool)\n                mask_padder[\n                    zyx_start[0]:zyx_start[0] + shape[0],\n                    zyx_start[1]:zyx_start[1] + shape[1],\n                    zyx_start[2]:zyx_start[2] + shape[2]\n                ] = result[key] if key != \"containers\" else result[key] + 0.5\n                result[key] = mask_padder\n\n    return result\ndef one_hot(label):\n    return np.stack( [label==i for i in range(6)] ,  0)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:13.515845Z","iopub.status.busy":"2025-02-01T19:13:13.515049Z","iopub.status.idle":"2025-02-01T19:13:13.548749Z","shell.execute_reply":"2025-02-01T19:13:13.547932Z"},"papermill":{"duration":0.042104,"end_time":"2025-02-01T19:13:13.549960","exception":false,"start_time":"2025-02-01T19:13:13.507856","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0500b864","cell_type":"code","source":"train_data = {}\nfor exp_name in tqdm(TRAIN_EXP):\n    train_data[exp_name] = {}\n    train_data[exp_name][\"volume\"] = read_one_data(exp_name, static_dir=f'{TRAIN_DIR}/static/ExperimentRuns')\n    train_data[exp_name][\"truth\"] = read_one_truth(exp_name, overlay_dir=f'{TRAIN_DIR}/overlay/ExperimentRuns')\n    pmin,pmax = np.percentile(train_data[exp_name][\"volume\"],(5,99))\n    train_data[exp_name][\"min\"] = pmin\n    train_data[exp_name][\"max\"] = pmax\n    \n    train_data[exp_name][\"label\"] = np.zeros((184,630,630) , dtype = np.int8)\n    train_data[exp_name][\"centroid\"] = np.zeros((184,630,630) , dtype = np.int8)\n    train_data[exp_name][\"containers\"] = np.zeros((184,630,630) , dtype = np.bool_)\n    #train_data[exp_name][\"heat_map\"] = np.zeros((184,630,630) , dtype = np.float16)\n    \n    for particle in train_data[exp_name][\"truth\"].keys():\n        radius = OBJECT_DICT[particle][\"radius\"]\n        radius_factor = np.log2(radius)/radius *.8\n        \n        label = OBJECT_DICT[particle][\"label\"]\n        \n        #train_data[exp_name][\"heat_map\"] += generate_heatmap((184,630,630) , train_data[exp_name][\"truth\"][particle])\n        \n        for point in train_data[exp_name][\"truth\"][particle]:\n            train_data[exp_name][\"containers\"] = draw_cylinder_in_image_fast(train_data[exp_name][\"containers\"], point, radius, z_factor = 1.7, yx_factor = 1.2, value = True)\n            train_data[exp_name][\"label\"] = draw_sphere_in_image_fast(train_data[exp_name][\"label\"], point, radius, radius_factor = radius_factor, value = label)\n            train_data[exp_name][\"centroid\"] = draw_dentroid(train_data[exp_name][\"centroid\"], point , value = label)\n    #train_data[exp_name][\"heat_map\"] = np.stack([train_data[exp_name][\"heat_map\"][particle] for particle in OBJECT_DICT.keys()], 0).astype(np.float16)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:13.562892Z","iopub.status.busy":"2025-02-01T19:13:13.562644Z","iopub.status.idle":"2025-02-01T19:13:40.047885Z","shell.execute_reply":"2025-02-01T19:13:40.046811Z"},"papermill":{"duration":26.493372,"end_time":"2025-02-01T19:13:40.049452","exception":false,"start_time":"2025-02-01T19:13:13.556080","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4e374c3e","cell_type":"code","source":"min_ = 0\nmax_ = 0\n\nfor k in train_data.keys():\n    pmin,pmax = np.percentile(train_data[k][\"volume\"],(5,99))\n    print(pmin,pmax)\n    min_ += pmin/7\n    max_ += pmax/7\n\nprint(min_,max_)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:40.064747Z","iopub.status.busy":"2025-02-01T19:13:40.064489Z","iopub.status.idle":"2025-02-01T19:13:49.591128Z","shell.execute_reply":"2025-02-01T19:13:49.590229Z"},"papermill":{"duration":9.535386,"end_time":"2025-02-01T19:13:49.592559","exception":false,"start_time":"2025-02-01T19:13:40.057173","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ba6f36d9","cell_type":"code","source":"for k in train_data.keys():\n    train_data[k][\"volume\"] = (train_data[k][\"volume\"]-min_)/(max_-min_)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:49.607932Z","iopub.status.busy":"2025-02-01T19:13:49.607702Z","iopub.status.idle":"2025-02-01T19:13:50.440512Z","shell.execute_reply":"2025-02-01T19:13:50.439812Z"},"papermill":{"duration":0.842257,"end_time":"2025-02-01T19:13:50.442188","exception":false,"start_time":"2025-02-01T19:13:49.599931","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"462af265","cell_type":"code","source":"def mean_std_shift (image,shift = 0.03):\n    factor = 1/(shift*2)\n    std = image.std()\n    mean = image.mean()\n    shift_mean = (torch.rand(1)/factor - shift).item()\n    shift_std = (torch.rand(1)/factor - shift).item()\n    new_mean = mean + mean * shift_mean\n    new_std = std + std * shift_std\n\n    new_image = (image-mean)/std*new_std+new_mean\n    return new_image\n\ndef generate_random_mask(mask_ratio, image_size):\n    \"\"\"\n    Generate a random numpy mask for a 3D image with a given mask ratio.\n\n    Parameters:\n        image_size (tuple): A tuple representing the dimensions of the 3D image (depth, height, width).\n        mask_ratio (float): A float between 0 and 1 indicating the fraction of elements to be masked (1 = fully masked, 0 = no mask).\n\n    Returns:\n        np.ndarray: A 3D numpy array with the same shape as the image, containing 0s (masked) and 1s (unmasked).\n    \"\"\"\n    if not (0 <= mask_ratio <= 1):\n        raise ValueError(\"mask_ratio must be between 0 and 1.\")\n\n    # Total number of elements in the image\n    total_elements = np.prod(image_size)\n\n    # Number of elements to mask\n    num_masked_elements = int(total_elements * mask_ratio)\n\n    # Create a flattened array with the specified number of masked (0) and unmasked (1) elements\n    mask_flat = np.ones(total_elements, dtype=np.uint8)\n    mask_flat[:num_masked_elements] = 0\n\n    # Shuffle the array to randomize the mask positions\n    np.random.shuffle(mask_flat)\n\n    # Reshape the flat mask array back to the original image size\n    mask = mask_flat.reshape(image_size)\n\n    return mask\n\n# Example usage\nimage_size = (96, 96, 96 )  # Example 3D image dimensions (depth, height, width)\nmask_ratio = 0.3           # 30% of the image will be masked\nmask = generate_random_mask(mask_ratio, image_size)\n\nprint((mask == 1).sum())","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.457706Z","iopub.status.busy":"2025-02-01T19:13:50.457447Z","iopub.status.idle":"2025-02-01T19:13:50.490925Z","shell.execute_reply":"2025-02-01T19:13:50.490115Z"},"papermill":{"duration":0.042403,"end_time":"2025-02-01T19:13:50.492161","exception":false,"start_time":"2025-02-01T19:13:50.449758","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"40172da7","cell_type":"code","source":"import torch.nn.functional as F","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.506832Z","iopub.status.busy":"2025-02-01T19:13:50.506613Z","iopub.status.idle":"2025-02-01T19:13:50.509543Z","shell.execute_reply":"2025-02-01T19:13:50.508921Z"},"papermill":{"duration":0.011519,"end_time":"2025-02-01T19:13:50.510777","exception":false,"start_time":"2025-02-01T19:13:50.499258","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"2d6df473","cell_type":"code","source":"class SegmentationDataset(Dataset):\n    def __init__(self, patch_size,length ,experiments = [\"TS_6_4\"]):\n        self.patch_size = patch_size\n        self.experiments = experiments\n        self.length = length\n        self.label_counts = {i:0 for i in range(6)}\n        for exp_name in experiments :\n            unique_values, counts = np.unique(train_data[exp_name][\"label\"], return_counts=True)\n            for i,unique in enumerate(unique_values):\n                self.label_counts[unique] += counts[i]\n        array = np.array([v for _,v in {0: 72996892, 1: 4176, 2: 1548, 3: 18574, 4: 6240, 5: 2170}.items()] )\n        array = 1/np.log1p(array)\n        self.weights = array/array.sum()\n        print(self.weights)\n    def __len__(self):\n        return self.length\n\n    def augment(self,result):\n        \n        do_flip_z = torch.rand(1)<.5\n        do_flip_y = torch.rand(1)<.5\n        do_flip_x = torch.rand(1)<.5\n        rot_times = random.choice([0,1,2,3])\n\n        for key in result.keys():\n            if do_flip_z:\n                result[key] = np.flip(result[key], axis=-3)\n            if do_flip_y:\n                result[key] = np.flip(result[key], axis=-2)\n            if do_flip_x:\n                result[key] = np.flip(result[key], axis=-1)\n\n            if rot_times !=0:\n                result[key] = np.rot90(result[key] , k = rot_times, axes=(-2,-1))\n        return result\n        \n    def _to_tensor(self,result):\n        for k in result.keys():\n            if \"label\" in k or \"heat_map\" in k:\n                #result[k] = numpy_one_hot(result[k])\n                result[k] = torch.tensor(result[k].copy() , dtype = torch.long if \"label\" in k else torch.float32)\n                #result[k] = one_hot_encode_3d(result[k] ,2)\n            else : \n                result[k] = torch.tensor(result[k].copy() , dtype = torch.float32)\n\n        #result[\"label\"] = torch.stack([result[k] for k in result.keys() if \"label\" in k])\n        return result\n        \n    def __getitem__(self,idx):\n        zyx = [random.choice(range(self.patch_size[i]//2 ,dim-self.patch_size[i]//2)) for i,dim in enumerate((184,630,630))]\n        exp_name = random.choice(self.experiments)\n        \n        result = {}\n        for key in [\"label\", \"containers\", \"volume\"]:#,\"heat_map\"\n            result [key] = crop_with_center(train_data[exp_name][key], self.patch_size, zyx)\n\n        result = self.augment(result)\n\n        result[\"volume\"] = mean_std_shift(result[\"volume\"])\n\n        result = pad(result,self.patch_size)\n\n        result[\"label_hot\"] = one_hot(result[\"label\"])\n\n        #result[\"label\"] = one_hot(result[\"label\"])\n        result = self._to_tensor(result)\n        return result\n","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.525321Z","iopub.status.busy":"2025-02-01T19:13:50.525099Z","iopub.status.idle":"2025-02-01T19:13:50.534978Z","shell.execute_reply":"2025-02-01T19:13:50.534328Z"},"papermill":{"duration":0.018412,"end_time":"2025-02-01T19:13:50.536200","exception":false,"start_time":"2025-02-01T19:13:50.517788","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8a6fe399","cell_type":"code","source":"import torch.nn as nn\n\nclass ConvBNReLU2D(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super(ConvBNReLU2D, self).__init__()\n        if kernel_size == 5:\n            padding = 2\n\n        if kernel_size == 7:\n            padding = 3\n            \n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\nclass ConvBNReLU3D(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super(ConvBNReLU3D, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.block(x)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.550580Z","iopub.status.busy":"2025-02-01T19:13:50.550337Z","iopub.status.idle":"2025-02-01T19:13:50.555617Z","shell.execute_reply":"2025-02-01T19:13:50.554938Z"},"papermill":{"duration":0.013749,"end_time":"2025-02-01T19:13:50.556826","exception":false,"start_time":"2025-02-01T19:13:50.543077","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e719f16e","cell_type":"code","source":"class dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.571162Z","iopub.status.busy":"2025-02-01T19:13:50.570964Z","iopub.status.idle":"2025-02-01T19:13:50.574140Z","shell.execute_reply":"2025-02-01T19:13:50.573567Z"},"papermill":{"duration":0.011517,"end_time":"2025-02-01T19:13:50.575308","exception":false,"start_time":"2025-02-01T19:13:50.563791","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ad0a2a03","cell_type":"code","source":"class Input(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3 , num_conv2d = 2):\n        super(Input, self).__init__()\n        self.norm = nn.BatchNorm2d(1)\n        self.conv2d_layers = nn.Sequential(\n            *[ConvBNReLU2D(in_channels * 4 if i == 0 else out_channels, out_channels) \n              for i in range(num_conv2d)]\n        )\n        self.conv2d_7 = ConvBNReLU2D(1,in_channels,7)\n        self.conv2d_5 = ConvBNReLU2D(1,in_channels,5)\n        self.conv2d_3 = ConvBNReLU2D(1,in_channels,3)\n\n    def forward(self, x):\n        # Apply Conv2D layers\n        b, c, d, h, w = x.shape  # Batch size, Channels, Depth, Height, Width\n        x = x.permute(0, 2, 1, 3, 4).reshape(b * d, c, h, w)  # Reshape for Conv2D\n        \n        x = self.norm(x)\n        x = torch.cat([x,self.conv2d_7(x), self.conv2d_5(x), self.conv2d_3(x)] , 1)\n        x = self.conv2d_layers(x)\n        # Reshape back to 5D for Conv3D layers\n        bd, c_out, h, w = x.shape\n        x = x.reshape(b, d, c_out, h, w).permute(0, 2, 1, 3, 4)\n        return dotdict({\"out\":x})\n\nimport torch.nn.functional as F\n\nclass EncoderBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, skip_channels = 0, num_conv3d=1, \n                 do_up = True , do_down = True ,use_transpose = False):\n        super(EncoderBlock, self).__init__()\n        self.do_up = do_up\n        \n        if self.do_up:\n            self.upsample = lambda x: F.interpolate(x, scale_factor=2, mode='trilinear')\n            \n        if use_transpose:\n            self.upsample = nn.Sequential(\n                nn.ConvTranspose3d(out_channels, out_channels, kernel_size=2, stride=2),\n                nn.BatchNorm3d(out_channels),\n                nn.ReLU(inplace=True)\n            )\n            \n        self.do_down = do_down \n        if self.do_down:\n            self.downsample = lambda x: F.interpolate(x, scale_factor=.5, mode='trilinear')\n\n        \n        self.conv3d_layers = nn.Sequential(\n            *[ConvBNReLU3D(out_channels if i!=0 else in_channels + skip_channels, out_channels, stride= (1, 1, 1)) \n              for i in range(num_conv3d)]\n        )\n\n    def forward(self, x , xskip = None):\n        if xskip is not None:\n            x = torch.cat([x, xskip], dim=1)\n\n        out = self.conv3d_layers(x)\n        output = {\n            \"out\": out,\n            \"up\":None,\n            \"down\":None\n        }\n\n        if self.do_up:\n            output[\"up\"] = self.upsample(out)   \n        if self.do_down:\n            output[\"down\"] = self.downsample(out)   \n        return dotdict(output)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.589963Z","iopub.status.busy":"2025-02-01T19:13:50.589765Z","iopub.status.idle":"2025-02-01T19:13:50.598768Z","shell.execute_reply":"2025-02-01T19:13:50.598154Z"},"papermill":{"duration":0.017748,"end_time":"2025-02-01T19:13:50.599997","exception":false,"start_time":"2025-02-01T19:13:50.582249","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"dcd8b420","cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self , channels = [28,32,36]):\n        super(Model, self).__init__()\n        self.register_buffer('D', torch.tensor(0))\n        self.output_type = ['particle', 'loss']\n\n        self.norm = nn.BatchNorm3d(1)\n        \n        self.encoder1 = EncoderBlock(in_channels = 1, out_channels = channels[0], num_conv3d=2 , do_up = False, do_down=True)\n        self.encoder2 = EncoderBlock(in_channels = channels[0], out_channels = channels[1], num_conv3d=2 , do_up = True, do_down=True)\n        \n        self.decoder1 = EncoderBlock(in_channels = channels[1], out_channels = channels[2], num_conv3d=4 , do_up = True, do_down=False)\n        self.decoder2 = EncoderBlock(in_channels = channels[2], out_channels = channels[1], num_conv3d=2 , skip_channels = channels[1], do_up = True, do_down=False , use_transpose = True)\n\n        self.pre = EncoderBlock(in_channels = channels[1], out_channels = channels[0], num_conv3d=2 , do_up = False, do_down=False)\n\n        self.mask = nn.Conv3d(channels[0], 6, 1, 1, bias=False)\n        \n        \n\n    def forward(self,batch):\n        device = self.D.device\n        volume = batch[\"volume\"].to(device).unsqueeze(1)\n\n        input_ = self.norm(volume)\n        \n        encode1 = self.encoder1(input_)\n        encode2 = self.encoder2(encode1.down)\n        \n        decode1 = self.decoder1(encode2.down)\n        #print(encode2.out.shape , decode1.up.shape)\n        decode2 = self.decoder2(encode2.out , decode1.up)\n\n        pre = self.pre(decode2.up)\n\n        logit = self.mask(pre.out)\n        #print(mask.shape)\n\n        output = {}\n        \n        if \"loss\" in self.output_type and \"label\" in batch.keys():\n        \n            # Apply weighted cross-entropy loss\n            output[\"loss\"] = F.cross_entropy(\n                logit, \n                batch['label'].to(device), \n                label_smoothing=0.01,\n            )\n\n        if \"particle\" in self.output_type:\n            output['particle'] = F.softmax(logit,1)\n            \n        return output\n","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.614706Z","iopub.status.busy":"2025-02-01T19:13:50.614485Z","iopub.status.idle":"2025-02-01T19:13:50.621596Z","shell.execute_reply":"2025-02-01T19:13:50.620960Z"},"papermill":{"duration":0.015979,"end_time":"2025-02-01T19:13:50.622813","exception":false,"start_time":"2025-02-01T19:13:50.606834","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"cef0ee3f","cell_type":"code","source":"def set_seed(seed):\n\n    random.seed(seed)  # Python's built-in random\n    np.random.seed(seed)  # NumPy random seed\n    torch.manual_seed(seed)  # Torch CPU random seed\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)  # Torch GPU random seed\n        torch.cuda.manual_seed_all(seed)  # All GPUs\n\n    # For deterministic behavior in CuDNN operations\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(80)","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.637234Z","iopub.status.busy":"2025-02-01T19:13:50.637016Z","iopub.status.idle":"2025-02-01T19:13:50.698285Z","shell.execute_reply":"2025-02-01T19:13:50.697630Z"},"papermill":{"duration":0.070173,"end_time":"2025-02-01T19:13:50.699777","exception":false,"start_time":"2025-02-01T19:13:50.629604","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a06b2e29","cell_type":"code","source":"def calculate_patch_starts(dimension_size: int, patch_size: int):\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) + 3\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    \nclass PredDataset(Dataset):\n    def __init__(self, patch_size, experiment = TRAIN_EXP[0] ,is_local = True):\n        self.is_local = is_local\n        self.experiment = experiment\n        self.patch_size = patch_size\n        #self.zyx = np.ones((3, 184+patch_size , 630+patch_size, 630+patch_size))*-1\n        pad_size = [patch_size[i]//2 for i in range(3)]\n        #self.zyx [:,pad_size:184+pad_size, pad_size:630+pad_size, pad_size:630+pad_size] = np.indices((184,630,630))\n        \n        self.volume = np.zeros((184+patch_size[0],630+patch_size[1],630+patch_size[2]))\n        self.volume [pad_size[0]:184+pad_size[0], pad_size[1]:630+pad_size[1], pad_size[2]:630+pad_size[2]] = train_data[experiment][\"volume\"]\n        self.label = np.zeros((184+patch_size[0],630+patch_size[1],630+patch_size[2]))\n        self.label [pad_size[0]:184+pad_size[0], pad_size[1]:630+pad_size[1], pad_size[2]:630+pad_size[2]] = train_data[experiment][\"label\"]\n        \n        self.locations = read_one_truth(experiment, overlay_dir=f'{TRAIN_DIR}/overlay/ExperimentRuns') if is_local else None\n\n        self.indexes = [[z,y,x] \n                       for z in calculate_patch_starts(184+patch_size[0],patch_size[0])\n                       for y in calculate_patch_starts(630+patch_size[1],patch_size[1])\n                       for x in calculate_patch_starts(630+patch_size[2],patch_size[2])]\n\n    def __len__(self):\n        return len(self.indexes)\n\n    def __getitem__(self,idx):\n        zyx = self.indexes [idx]\n        patch = self.volume[zyx[0]:zyx[0]+self.patch_size[0],zyx[1]:zyx[1]+self.patch_size[1],zyx[2]:zyx[2]+self.patch_size[2]]\n        label = self.label [zyx[0]:zyx[0]+self.patch_size[0],zyx[1]:zyx[1]+self.patch_size[1],zyx[2]:zyx[2]+self.patch_size[2]]\n        shape = patch.shape\n        if shape[0]*shape[1]*shape[2] != self.patch_size[0]*self.patch_size[1]*self.patch_size[2]:\n            padder = np.zeros(self.patch_size)\n            padder[:shape[0],:shape[1],:shape[2]] = patch\n            patch = padder\n            padder[:shape[0],:shape[1],:shape[2]] = label\n            label = padder\n            \n        return {\"volume\":torch.tensor(patch,dtype = torch.float32),'zyx':  torch.tensor(zyx,dtype = torch.long),\"label\":torch.tensor(label,dtype = torch.long) , \"label_hot\":torch.tensor(one_hot(label),dtype = torch.long)}\n\n\n\ndef get_probs():\n    val_loss = 0\n    \n    weight = torch.zeros((patch_size[0], patch_size[1], patch_size[2]) , dtype =torch.float16).to(\"cuda\")\n    weight[8:patch_size[0]-8, 8:patch_size[1]-8, 8:patch_size[2]-8] += 1\n    \n    # Initialize output tensors\n    all_logits = torch.zeros((6, 184 + patch_size[0], 630 + patch_size[1], 630 + patch_size[2]) , dtype =torch.float16).to(\"cuda\")\n    count = torch.zeros((184 + patch_size[0], 630 + patch_size[1], 630 + patch_size[2]) , dtype =torch.float16).to(\"cuda\")\n    \n    model.output_type = [\"particle\",\"loss\"]\n    model.eval()\n    with torch.no_grad():\n        with torch.amp.autocast(\"cuda\"):\n            for batch in tqdm(pl):\n                zyxs = batch[\"zyx\"]\n                outputs = model(batch)\n                local_logits = outputs[\"particle\"]\n                \n                val_loss += outputs[\"loss\"].item()\n                for i in range(len(zyxs)):\n                    zyx = zyxs[i]\n                    count[zyx[0]:zyx[0]+patch_size[0], zyx[1]:zyx[1]+patch_size[1], zyx[2]:zyx[2]+patch_size[2]] += weight\n        \n                    all_logits[:, zyx[0]:zyx[0]+patch_size[0], zyx[1]:zyx[1]+patch_size[1], zyx[2]:zyx[2]+patch_size[2]] += local_logits[i] * weight\n                    \n        # Print epoch loss\n        print(f\" Val Loss: {val_loss/len(pl):.8f}\")\n        # Crop to remove padding\n        all_logits = all_logits[:, patch_size[0]//2:patch_size[0]//2+184, patch_size[1]//2:patch_size[1]//2+630, patch_size[2]//2:patch_size[2]//2+630]\n        count = count[patch_size[0]//2:patch_size[0]//2+184, patch_size[1]//2:patch_size[1]//2+630, patch_size[2]//2:patch_size[2]//2+630]\n        probs = nn.Softmax(dim=0)(torch.tensor(all_logits/count)).detach().cpu().numpy()\n        # Compute probabilities\n        probs = (all_logits/count).detach().cpu().numpy()\n    \n        # Convert tensors to numpy arrays\n        all_logits = all_logits.detach().cpu().numpy()\n        count = count.detach().cpu().numpy()\n        weight = weight.detach().cpu().numpy()\n\n    return probs\n\ndef evaluate_predictions(stats, pred_loader, distance_threshold=3, beta=4 , particle_name = None):\n    best_f_beta = 0\n    best_metric = None\n    voxel_threshold = 5\n    # Filter predictions based on voxel count\n    pred = np.array([centroid for i, centroid in enumerate(stats[\"centroids\"]) if i != 0 and stats[\"voxel_counts\"][i] > voxel_threshold])\n    truth_locations = np.array(pred_loader.dataset.locations[particle_name])\n    # Perform evaluation\n    hit, fp, miss, metric = do_one_eval(truth_locations, pred, distance_threshold)\n\n    # Calculate precision, recall, and F-beta score\n    precision = len(hit[0]) / (len(hit[0]) + len(fp)) if (len(hit[0]) + len(fp)) > 0 else 0\n    recall = len(hit[0]) / (len(hit[0]) + len(miss)) if (len(hit[0]) + len(miss)) > 0 else 0\n\n    beta_squared = beta ** 2\n    f_beta = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (precision + recall) > 0 else 0\n    if f_beta>= best_f_beta:\n        best_f_beta = f_beta\n        best_metric = {\n            \"truth\": len(truth_locations),\n            \"predict\": len(pred),\n            \"hit\": len(hit[0]),\n            \"fp\": len(fp),\n            \"miss\": len(miss),\n            \"f_b\": f_beta,\n            \"thresh\": voxel_threshold\n        }\n    # Return results as JSON-like dictionary\n    return best_metric\n    \ndef do_one_eval(truth, predict, threshold = 3):\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\ndef do_evaluate():\n    probs = get_probs()\n    evals = {}\n    for particle in OBJECT_DICT.keys():\n        label = OBJECT_DICT[particle][\"label\"]\n        thresh = OBJECT_DICT[particle][\"radius\"]/2\n        labels_out = cc3d.connected_components(np.array(probs[label, :, :, :] > 0.06), connectivity=18)\n        stats = cc3d.statistics(labels_out)\n\n        evals[particle] = evaluate_predictions(stats, pl, distance_threshold = thresh, particle_name = particle)\n    return evals","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.715111Z","iopub.status.busy":"2025-02-01T19:13:50.714861Z","iopub.status.idle":"2025-02-01T19:13:50.739334Z","shell.execute_reply":"2025-02-01T19:13:50.738533Z"},"papermill":{"duration":0.033576,"end_time":"2025-02-01T19:13:50.740714","exception":false,"start_time":"2025-02-01T19:13:50.707138","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d6028dfe","cell_type":"code","source":"","metadata":{"papermill":{"duration":0.006606,"end_time":"2025-02-01T19:13:50.754292","exception":false,"start_time":"2025-02-01T19:13:50.747686","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"6e614fa0","cell_type":"code","source":"import torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.nn.utils import clip_grad_norm_\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport optuna\n\n\ndef find_best_learning_rate(model, objective_loader , learning_rates):\n    \"\"\"Find the best learning rate using Optuna.\"\"\"\n    def objective(trial):\n        lr = trial.suggest_loguniform(\"lr\", learning_rates[1], learning_rates[0])  # Suggest a learning rate in a wide range\n        optimizer = optim.Adam(model.parameters(), lr=lr, betas=(0.9, 0.999))\n        scaler = GradScaler()\n        \n        model.train()\n        total_loss = 0.0\n        for batch in tqdm(objective_loader):\n            with autocast():\n                outputs = model(batch)\n                loss = outputs[\"loss\"]\n\n            optimizer.zero_grad()\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            total_loss += loss.item()\n        return total_loss / len(objective_loader)\n\n    study = optuna.create_study(direction=\"minimize\")\n    study.optimize(objective, n_trials=4)  # Adjust n_trials as needed for computational limits\n\n    best_lr = study.best_params[\"lr\"]\n    print(f\"Best Learning Rate: {best_lr}\")\n    return best_lr\n\nlearning_rates = {\n    0 : [1e-3 , 1e-4],\n    10 : [1e-4 , 1e-5],\n    20 : [1e-5 , 1e-7],\n    25 : [1e-5 , 1e-7],\n    30 : [1e-6 , 1e-8],\n    35 : [1e-7 , 1e-10],\n}\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Hyperparameters\nbatch_size = 4\nnum_epochs = 60\npatch_size = (128, 128, 128)\n\nfolds = []\ntrain_experiments = [TRAIN_EXP[i] for i in range(7) if i not in folds]\n\n# Functions\ndef get_model():\n    \"\"\"Initialize and return the model.\"\"\"\n    model = Model().to(device)\n    return model\n\ndef get_loaders():\n    \"\"\"Prepare and return the dataloaders.\"\"\"\n    train_dataset = SegmentationDataset(patch_size, 1024, experiments=train_experiments)\n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\n    \n    objective_dataset = SegmentationDataset(patch_size, 1024, experiments=train_experiments)\n    objective_loader = DataLoader(objective_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\n    \n    return train_loader, objective_loader\n\ndef train_one_epoch(model, train_loader, optimizer, scaler, max_norm=1.0):\n    \"\"\"Train the model for one epoch with gradient clipping and return the average loss.\"\"\"\n    model.train()\n    model.output_type = [\"loss\"]\n    train_loss = 0.0\n\n    for batch in tqdm(train_loader, desc=\"Training\", leave=False):\n        with autocast():\n            outputs = model(batch)\n            loss = outputs[\"loss\"]\n\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n\n        # Unscales the gradients of optimizer's assigned params and clips gradients\n        scaler.unscale_(optimizer)  \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)\n\n        scaler.step(optimizer)\n        scaler.update()\n\n        train_loss += loss.item()\n\n    avg_train_loss = train_loss / len(train_loader)\n    return avg_train_loss\n\ndef train_model():\n    \"\"\"Train the model across all epochs.\"\"\"\n    model = get_model()\n    train_loader , objective_loader = get_loaders()\n\n    scaler = GradScaler()\n    learning_rate = 1e-4\n    for epoch in range(num_epochs):\n\n\n        if epoch == 35 :        \n            learning_rate = 1e-5\n\n        if epoch == 45 :        \n            learning_rate = 1e-7\n            \n        optimizer = optim.Adam(\n            model.parameters(), \n            lr=learning_rate, \n            betas=(0.9, 0.999)\n        )\n        \n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        avg_train_loss = train_one_epoch(model, train_loader, optimizer, scaler)\n        print(f\"Train Loss: {avg_train_loss:.4f}\")\n\n        # Save model checkpoint\n        checkpoint_path = f\"model_all_{epoch}.bin\"\n        torch.save(model.state_dict(), checkpoint_path)\n\n    # Free GPU memory\n    torch.cuda.empty_cache()\n\n# Call the training function\ntrain_model()","metadata":{"execution":{"iopub.execute_input":"2025-02-01T19:13:50.768912Z","iopub.status.busy":"2025-02-01T19:13:50.768696Z","iopub.status.idle":"2025-02-02T03:55:51.050399Z","shell.execute_reply":"2025-02-02T03:55:51.049251Z"},"papermill":{"duration":31320.991515,"end_time":"2025-02-02T03:55:51.752840","exception":false,"start_time":"2025-02-01T19:13:50.761325","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b1973286","cell_type":"code","source":"","metadata":{"papermill":{"duration":0.703855,"end_time":"2025-02-02T03:55:53.284530","exception":false,"start_time":"2025-02-02T03:55:52.580675","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}