{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n### **Apo-Ferritin Identification Using a Simple 3D U-Net Model**\n\nTesting identification of apo-ferritin complexes in tomograms using a simple 3D U-Net model. \n\nThe key steps include data preparation, augmentation, model training, and evaluation. \n\n#### **Overview**\n1. **Load Tomograms**:\n   - Load tomograms from 7 experimental datasets.\n   - Use the second resolution (92x315x315) for computational efficiency without significant loss of information.\n\n2. **Data Augmentation**:\n   - Address the issue of insufficient labels (center coordinates).\n   - Identify additional low-intensity regions in the tomograms to augment the labels.\n   - Create additional apo-ferritin labels by copying the existing complexes to new coordinates.\n\n3. **Build Binary Masks**:\n   - Generate binary segmentation masks for apo-ferritin using center coordinates and radius to define the complex regions.\n\n4. **Model Design and Training**:\n   - Implement a simple 3D U-Net model for apo-ferritin segmentation.\n   - Use weighted Dice loss to handle class imbalance and optimize model performance.\n   - Train the model on augmented datasets.\n\n5. **Model Evaluation**:\n   - Evaluate the model on validation datasets using key metrics such as Dice Coefficient, IoU, Precision, and Recall.\n   - Resize predicted masks to match ground truth dimensions for accurate evaluation.\n\n6. **Prediction**:\n   - Predict and analyze the results for test datasets.\n","metadata":{}},{"cell_type":"code","source":"#!pip install zarr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-31T00:40:53.540729Z","iopub.execute_input":"2025-01-31T00:40:53.541052Z","iopub.status.idle":"2025-01-31T00:40:53.544983Z","shell.execute_reply.started":"2025-01-31T00:40:53.541021Z","shell.execute_reply":"2025-01-31T00:40:53.544026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = {\n    \"name\": \"czii_cryoet_mlchallenge_2024\",\n    \"description\": \"2024 CZII CryoET ML Challenge training data.\",\n    \"version\": \"1.0.0\",\n\n    \"pickable_objects\": [\n        {\n            \"name\": \"apo-ferritin\",\n            \"is_particle\": True,\n            \"pdb_id\": \"4V1W\",\n            \"label\": 1,\n            \"color\": [  0, 117, 220, 128],\n            \"radius\": 60,\n            \"map_threshold\": 0.0418\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/test/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": True\n    },\n\n    \"static_root\": \"/kaggle/input/czii-cryo-et-object-identification/test/static\"\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:45.310330Z","iopub.execute_input":"2025-01-30T23:08:45.310674Z","iopub.status.idle":"2025-01-30T23:08:45.316188Z","shell.execute_reply.started":"2025-01-30T23:08:45.310641Z","shell.execute_reply":"2025-01-30T23:08:45.315235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport zarr\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\nimport pandas as pd\n\nfrom datetime import datetime","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:45.317167Z","iopub.execute_input":"2025-01-30T23:08:45.317487Z","iopub.status.idle":"2025-01-30T23:08:45.750363Z","shell.execute_reply.started":"2025-01-30T23:08:45.317455Z","shell.execute_reply":"2025-01-30T23:08:45.749561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load tomogram data\nbase_path = \"/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns\"\noverlay_base_path = \"/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns/\"\nexperiment_mapping = {\n    1: \"TS_5_4\",\n    2: \"TS_6_4\",\n    3: \"TS_6_6\",\n    4: \"TS_69_2\",\n    5: \"TS_73_6\",\n    6: \"TS_86_3\",\n    7: \"TS_99_9\"\n}\n\ntomograms = {}\nfor dataset_id, experiment in experiment_mapping.items():\n    zarr_file_path = os.path.join(base_path, experiment, \"VoxelSpacing10.000/denoised.zarr\")\n    if os.path.exists(zarr_file_path):\n        print(f\"Loading dataset {dataset_id}: {zarr_file_path}\")\n        tomograms[dataset_id] = zarr.open(zarr_file_path, mode='r')\n    else:\n        print(f\"Tomogram for {experiment} not found at {zarr_file_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:45.752055Z","iopub.execute_input":"2025-01-30T23:08:45.752498Z","iopub.status.idle":"2025-01-30T23:08:45.925602Z","shell.execute_reply.started":"2025-01-30T23:08:45.752475Z","shell.execute_reply":"2025-01-30T23:08:45.924956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#check the resolutions on temograms\nfor dataset_id, tomogram in tomograms.items():\n    print(f\"Dataset ID: {dataset_id}\")\n    print(tomogram.tree())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:45.926630Z","iopub.execute_input":"2025-01-30T23:08:45.926869Z","iopub.status.idle":"2025-01-30T23:08:46.132806Z","shell.execute_reply.started":"2025-01-30T23:08:45.926848Z","shell.execute_reply":"2025-01-30T23:08:46.132206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Particle properties \nvoxel_spacing = 10.0\nresolution = 2 # using 2nd resolution which is 92x315x315\napo_ferritin = config['pickable_objects'][0]\nvoxel_radius = int(apo_ferritin['radius'] / voxel_spacing / resolution)\nthreshold = apo_ferritin['map_threshold']\nprint(f\"radius: {voxel_radius}, threshold: {threshold}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:46.133495Z","iopub.execute_input":"2025-01-30T23:08:46.133685Z","iopub.status.idle":"2025-01-30T23:08:46.138506Z","shell.execute_reply.started":"2025-01-30T23:08:46.133667Z","shell.execute_reply":"2025-01-30T23:08:46.137784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to load particle centers from a JSON file and convert to voxel indices\ndef load_centers(json_file_path, voxel_spacing):\n    try:\n        with open(json_file_path, 'r') as f:\n            particle_locations = json.load(f)\n        \n        particle_points = particle_locations.get(\"points\", [])\n        if not particle_points:\n            print(f\"No 'points' found in JSON file: {json_file_path}\")\n            return np.array([])\n\n        # Extract 3D coordinates and convert to voxel indices\n        voxel_coordinates = [\n            (\n                round(point[\"location\"][\"z\"] / voxel_spacing),  # z-index\n                round(point[\"location\"][\"y\"] / voxel_spacing),  # y-index\n                round(point[\"location\"][\"x\"] / voxel_spacing)   # x-index\n            )\n            for point in particle_points\n        ]\n\n        return np.array(voxel_coordinates)\n\n    except Exception as e:\n        print(f\"Error loading JSON file {json_file_path}: {e}\")\n        return np.array([])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:46.139111Z","iopub.execute_input":"2025-01-30T23:08:46.139339Z","iopub.status.idle":"2025-01-30T23:08:46.152960Z","shell.execute_reply.started":"2025-01-30T23:08:46.139294Z","shell.execute_reply":"2025-01-30T23:08:46.152348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# visualize 1 apo-ferritin complex from each experiment on 2nd resolution (92x315x315) \n\n# Display slice with and without protein complex highlighter\ndef display_slice_with_and_without_marks(data, center, radius):\n    z, y, x = center  \n    slice_index = z  # displaying the z-plane of the center\n\n    # Ensure the slice index is within bounds\n    if 0 <= slice_index < data.shape[0]:\n        fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n\n        # Without marks\n        axes[0].imshow(data[slice_index, :, :], cmap='gray', origin='lower')\n        axes[0].set_title(f\"Slice {slice_index} Without Marks\")\n        axes[0].set_xlabel(\"X\")\n        axes[0].set_ylabel(\"Y\")\n\n        # With marks\n        axes[1].imshow(data[slice_index, :, :], cmap='gray', origin='lower')\n        #using radius+3 to celarly see protein complex inside the circle\n        circle = plt.Circle((x, y), radius+3, color='red', fill=False, label='apo-ferritin')\n        axes[1].add_artist(circle)\n        axes[1].set_title(f\"Slice {slice_index} With Center and Radius\")\n        axes[1].set_xlabel(\"X\")\n        axes[1].set_ylabel(\"Y\")\n        axes[1].legend()\n\n        plt.tight_layout()\n        plt.show()\n    else:\n        print(f\"Slice index {slice_index} is out of bounds for the data shape {data.shape}.\")\n        \n\nfor exp_id in range(1, 8):\n    experiment = experiment_mapping[exp_id]\n    \n    try:\n        # Access second-resolution data ('1' level)\n        tomogram = tomograms[exp_id]\n        data = tomogram['1'][:]  # Access the second-resolution array\n        \n        # Load the center coordinates for the dataset\n        json_file_path = os.path.join(overlay_base_path, experiment, \"Picks\", \"apo-ferritin.json\")\n        centers = load_centers(json_file_path, voxel_spacing)\n        print(f\"{experiment} number of label centers: {len(centers)}\")\n    \n        if centers.size != 0:\n            # Adjust center coordinates for resolution\n            adjusted_centers = [(z // resolution, y // resolution, x // resolution) for z, y, x in centers]\n    \n            # Limit to the first 5 centers\n            for center in adjusted_centers[:1]:\n                print(f\"center: {center}\")\n                x, y, z = center  # Center in voxel indices\n                # Display the slice with and without marks\n                display_slice_with_and_without_marks(data, center, voxel_radius)\n    \n    except KeyError as e:\n        print(f\"KeyError: {e}. Skipping Dataset ID: {exp_id}\")\n    except Exception as e:\n        print(f\"Error processing Dataset ID {exp_id}: {e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:46.153617Z","iopub.execute_input":"2025-01-30T23:08:46.153812Z","iopub.status.idle":"2025-01-30T23:08:54.483932Z","shell.execute_reply.started":"2025-01-30T23:08:46.153796Z","shell.execute_reply":"2025-01-30T23:08:54.483120Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"***We can observe a small ring structure inside each red circle, which we assume to be apo-ferritin.***","metadata":{}},{"cell_type":"code","source":"# verify the gray intensity stats of the protein complex regions in tomogram\n\ndef calculate_intensity_stats(region):\n    #Calculate intensity statistics for a given region.\n    return {\n        \"mean\": np.mean(region),\n        \"median\": np.median(region),\n        \"min\": np.min(region),\n        \"max\": np.max(region),\n        \"std\": np.std(region),\n    }\n\ndef generate_intensity_stats_dataframe(data, centers, voxel_radius):\n    #Generate a DataFrame containing intensity statistics for all centers.\n    stats_list = []\n    for i, center in enumerate(centers, start=1):\n        x, y, z = center\n        # Extract region based on radius\n        region = data[\n            max(0, x-voxel_radius):min(data.shape[0], x+voxel_radius+1),\n            max(0, y-voxel_radius):min(data.shape[1], y+voxel_radius+1),\n            max(0, z-voxel_radius):min(data.shape[2], z+voxel_radius+1)\n        ]\n        # Calculate intensity statistics\n        stats = calculate_intensity_stats(region)\n        stats[\"center\"] = center\n        stats[\"particle_id\"] = i\n        stats_list.append(stats)\n    \n    # Convert the list of stats into a DataFrame\n    return pd.DataFrame(stats_list)\n\n\nintensity_stats_df = generate_intensity_stats_dataframe(data, adjusted_centers, voxel_radius)\n\nintensity_stats_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:54.484909Z","iopub.execute_input":"2025-01-30T23:08:54.485213Z","iopub.status.idle":"2025-01-30T23:08:54.531348Z","shell.execute_reply.started":"2025-01-30T23:08:54.485181Z","shell.execute_reply":"2025-01-30T23:08:54.530738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#identify low intensity segments for augmenting additional complexes (labels).\n#A maximum intensity threshold of 6.0E-06 was used to identify regions \n\nfrom joblib import Parallel, delayed\nfrom scipy.spatial import KDTree\nfrom tqdm import tqdm\n\n#Identify low-intensity spherical segments in the tomogram using parallel processing.\ndef find_low_intensity_segments(data, existing_centers, voxel_radius, intensity_threshold, max_segments, n_jobs=16):\n       \n    checked_voxels = np.zeros_like(data, dtype=bool)\n\n    # Convert existing centers into a mask of already occupied regions\n    for center in existing_centers:\n        x, y, z = center\n        xx, yy, zz = np.meshgrid(\n            np.arange(-voxel_radius, voxel_radius + 1),\n            np.arange(-voxel_radius, voxel_radius + 1),\n            np.arange(-voxel_radius, voxel_radius + 1),\n            indexing=\"ij\",\n        )\n        distances = np.sqrt(xx**2 + yy**2 + zz**2)\n        sphere_mask = distances <= voxel_radius\n\n        # Apply mask to avoid overlaps\n        cx, cy, cz = np.meshgrid(\n            np.clip(np.arange(x - voxel_radius, x + voxel_radius + 1), 0, data.shape[0] - 1),\n            np.clip(np.arange(y - voxel_radius, y + voxel_radius + 1), 0, data.shape[1] - 1),\n            np.clip(np.arange(z - voxel_radius, z + voxel_radius + 1), 0, data.shape[2] - 1),\n            indexing=\"ij\",\n        )\n        checked_voxels[cx, cy, cz] |= sphere_mask\n    \n    #print(datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\": Completed exiting centers masks conversion...\") \n        \n    def process_chunk(x_range):\n        #print(datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\": execution started for chunk {x_range}...\", flush=True) \n        local_segments = []\n        for x in x_range:\n            for y in range(voxel_radius, data.shape[1] - voxel_radius):\n                for z in range(voxel_radius, data.shape[2] - voxel_radius):\n                    if checked_voxels[x, y, z]:\n                        continue  # Skip already occupied regions\n\n                    # Extract spherical region around the candidate center\n                    xx, yy, zz = np.meshgrid(\n                        np.arange(-voxel_radius, voxel_radius + 1),\n                        np.arange(-voxel_radius, voxel_radius + 1),\n                        np.arange(-voxel_radius, voxel_radius + 1),\n                        indexing=\"ij\",\n                    )\n                    distances = np.sqrt(xx**2 + yy**2 + zz**2)\n                    sphere_mask = distances <= voxel_radius\n\n                    region = data[\n                        max(0, x - voxel_radius):min(data.shape[0], x + voxel_radius + 1),\n                        max(0, y - voxel_radius):min(data.shape[1], y + voxel_radius + 1),\n                        max(0, z - voxel_radius):min(data.shape[2], z + voxel_radius + 1),\n                    ]\n\n                    if np.all(region[sphere_mask] < intensity_threshold):\n                        local_segments.append((x, y, z))\n\n                    if len(local_segments) >= max_segments:\n                        return local_segments\n        return local_segments\n\n    # Split the x-axis into chunks for parallel processing\n    x_ranges = np.array_split(range(voxel_radius, data.shape[0] - voxel_radius), n_jobs)\n    results = Parallel(n_jobs=n_jobs)(delayed(process_chunk)(x_range) for x_range in x_ranges)\n\n    # Combine results from all chunks\n    new_segments = [seg for result in results for seg in result]\n\n    return new_segments\n    \n#Remove overlapping segments using KD-Tree for fast nearest-neighbor queries.\ndef remove_overlapping_segments(segments, voxel_radius):\n    unique_segments = []\n    tree = None\n\n    for seg in tqdm(segments, desc=\"Removing Overlapping Segments\"):\n        if tree is None:\n            unique_segments.append(seg)\n            tree = KDTree(unique_segments)\n            continue\n\n        # Check if the segment is too close to any existing segment\n        distances, indices = tree.query([seg], k=1)\n        if distances[0] > voxel_radius:\n            unique_segments.append(seg)\n            tree = KDTree(unique_segments)  # Rebuild tree with new segment\n\n    return unique_segments\n\n\nintensity_threshold = 6.0E-06\nmax_segments = 200000\n# We will use 5 experiments for training\ntarget_segments = {}\nfor exp_id in range(1, 6):\n    experiment = experiment_mapping[exp_id]\n    tomogram = tomograms[exp_id]\n    data = tomogram['1'][:]  # Access the second-resolution array\n    \n    # Load the center coordinates for the dataset\n    json_file_path = os.path.join(overlay_base_path, experiment, \"Picks\", \"apo-ferritin.json\")\n    centers = load_centers(json_file_path, voxel_spacing)\n    adjusted_centers = [(z // resolution, y // resolution, x // resolution) for z, y, x in centers]\n    \n    print(\"\\n\",datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\" : {experiment} : execution begins...\")        \n    new_segments = find_low_intensity_segments(data, adjusted_centers, voxel_radius, intensity_threshold, max_segments, n_jobs=16)\n    print(datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\" : {experiment} : Received {len(new_segments)} low-intensity segments.\")\n    \n    # Filter out overlapping segments\n    filtered_segments = remove_overlapping_segments(new_segments, voxel_radius)\n    target_segments[exp_id] = filtered_segments\n    print(f\"{experiment} : Identified {len(filtered_segments)} target segments.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:08:54.532118Z","iopub.execute_input":"2025-01-30T23:08:54.532354Z","iopub.status.idle":"2025-01-30T23:39:07.038142Z","shell.execute_reply.started":"2025-01-30T23:08:54.532333Z","shell.execute_reply":"2025-01-30T23:39:07.037244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for key, value in target_segments.items():\n    print(f\"Key: {key}, Length: {len(value)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:07.039137Z","iopub.execute_input":"2025-01-30T23:39:07.039983Z","iopub.status.idle":"2025-01-30T23:39:07.044707Z","shell.execute_reply.started":"2025-01-30T23:39:07.039952Z","shell.execute_reply":"2025-01-30T23:39:07.043987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#random selection of target_segments to ensure labels volume variability \nimport random\n# divisors for each key\ndivisors = {1: 1, 2: 5, 3: 1, 4: 10, 5: 20}\n\n# Select random values based on the divisor\nrandom_target_segments = {\n    key: random.sample(value, len(value) // divisors[key])\n    for key, value in target_segments.items()\n}\n\n# Print the length of the selected random samples\nfor key, value in random_target_segments.items():\n    print(f\"Key: {key}, Selected Length: {len(value)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:07.045519Z","iopub.execute_input":"2025-01-30T23:39:07.045794Z","iopub.status.idle":"2025-01-30T23:39:07.078942Z","shell.execute_reply.started":"2025-01-30T23:39:07.045773Z","shell.execute_reply":"2025-01-30T23:39:07.078288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Augment training data by copying and transforming apo-ferritin particles.\ndef copy_and_transform_particles(data, centers, to_segments, voxel_radius):\n    from datetime import datetime\n    from tqdm import tqdm\n\n    augmented_data = np.copy(data)  # Create a copy to avoid modifying the original\n    sphere_size = 2 * voxel_radius + 1\n    \n    successful_targets = list(centers)  # Track successful target centers\n\n    for target_center in tqdm(to_segments, desc='Processing Target Centers'):\n        # Select a random source center\n        source_center = random.choice(centers)\n        sx, sy, sz = source_center\n        tx, ty, tz = target_center\n\n        try:\n            # Extract the source spherical region\n            source_region = data[\n                max(0, sx - voxel_radius):min(data.shape[0], sx + voxel_radius + 1),\n                max(0, sy - voxel_radius):min(data.shape[1], sy + voxel_radius + 1),\n                max(0, sz - voxel_radius):min(data.shape[2], sz + voxel_radius + 1)\n            ]\n\n            # Generate a sphere mask\n            xx, yy, zz = np.meshgrid(\n                np.arange(-voxel_radius, voxel_radius + 1),\n                np.arange(-voxel_radius, voxel_radius + 1),\n                np.arange(-voxel_radius, voxel_radius + 1),\n                indexing=\"ij\"\n            )\n            distances = np.sqrt(xx**2 + yy**2 + zz**2)\n            sphere_mask = distances <= voxel_radius\n\n            # Adjust the mask to match the source region shape\n            adjusted_mask = sphere_mask[\n                :source_region.shape[0],\n                :source_region.shape[1],\n                :source_region.shape[2]\n            ]\n\n            if source_region.shape != adjusted_mask.shape:\n                print(f\"Skipping source center {source_center}: Shape mismatch.\")\n                continue\n                \n            if (sx - voxel_radius < 0 or sy - voxel_radius < 0 or sz - voxel_radius < 0 or\n            sx + voxel_radius >= data.shape[0] or sy + voxel_radius >= data.shape[1] or sz + voxel_radius >= data.shape[2]):\n                #print(f\"Skipping source center {source_center}: Too close to boundary.\")\n                continue\n\n            if (tx - voxel_radius < 0 or ty - voxel_radius < 0 or tz - voxel_radius < 0 or\n            tx + voxel_radius >= data.shape[0] or ty + voxel_radius >= data.shape[1] or tz + voxel_radius >= data.shape[2]):\n                #print(f\"Skipping target center {target_center}: Too close to boundary.\")\n                continue\n\n\n            # Apply transformations (random rotation or flipping)\n            source_region_transformed = np.copy(source_region)\n            if random.choice([True, False]) and source_region.ndim == 3:\n                source_region_transformed = np.flip(source_region_transformed, axis=random.choice([0, 1, 2]))\n            if random.choice([True, False]) and source_region.ndim == 3:\n                source_region_transformed = np.rot90(source_region_transformed, k=random.randint(1, 3), axes=(0, 1))\n\n            # Extract the target region\n            target_region = augmented_data[\n                max(0, tx - voxel_radius):min(data.shape[0], tx + voxel_radius + 1),\n                max(0, ty - voxel_radius):min(data.shape[1], ty + voxel_radius + 1),\n                max(0, tz - voxel_radius):min(data.shape[2], tz + voxel_radius + 1)\n            ]\n\n            # Adjust the mask to match the target region shape\n            adjusted_mask = sphere_mask[\n                :target_region.shape[0],\n                :target_region.shape[1],\n                :target_region.shape[2]\n            ]\n\n            if target_region.shape != adjusted_mask.shape:\n                print(f\"Skipping target center {target_center}: Target region shape mismatch.\")\n                continue\n\n            # Apply the transformed region\n            target_region[adjusted_mask] = source_region_transformed[adjusted_mask]\n            \n            # Add to successful targets\n            if target_center not in successful_targets:\n                successful_targets.append(target_center)\n                \n        except Exception as e:\n            print(f\"Error at target center {target_center}: {e}\")\n            print(f\"Source Center: {source_center}, Target Center: {target_center}\")\n            print(f\"Source Region Shape: {source_region.shape if 'source_region' in locals() else 'Unavailable'}\")\n            print(f\"Target Region Shape: {target_region.shape if 'target_region' in locals() else 'Unavailable'}\")\n\n    \n    print(\"Copied targets \", len(successful_targets)-len(centers), \" out of \", len(to_segments))\n    print(\"total targets: \", len(successful_targets))\n    \n    return augmented_data, successful_targets\n\n\naugmented_tomograms = {}\nfinal_targets = {}\nfor exp_id in range(1, 6):\n    experiment = experiment_mapping[exp_id]\n    tomogram = tomograms[exp_id]\n    data = tomogram['1'][:]  # Access the second-resolution array\n    \n    # Load the center coordinates for the dataset\n    json_file_path = os.path.join(overlay_base_path, experiment, \"Picks\", \"apo-ferritin.json\")\n    centers = load_centers(json_file_path, voxel_spacing)\n    adjusted_centers = [(z // resolution, y // resolution, x // resolution) for z, y, x in centers]\n    to_segments = random_target_segments[exp_id]\n    print(\"\\n\",datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\" : {experiment} : execution begins...\")        \n    augmented_data, augmented_targets = copy_and_transform_particles(data, adjusted_centers, to_segments, voxel_radius)\n    augmented_tomograms[exp_id] = augmented_data\n    final_targets[exp_id] = augmented_targets\n    print(datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\"), f\"Augmentation completed. Data shape: {augmented_data.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:07.081849Z","iopub.execute_input":"2025-01-30T23:39:07.082087Z","iopub.status.idle":"2025-01-30T23:39:19.296081Z","shell.execute_reply.started":"2025-01-30T23:39:07.082065Z","shell.execute_reply":"2025-01-30T23:39:19.295350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Visualize augmented apo-ferritin in circular regions.\ndef visualize_apo_ferritin(data, original_data, centers, radius, title):\n    \n    fig, axes = plt.subplots(len(centers), 3, figsize=(15, len(centers) * 5))\n\n    for i, center in enumerate(centers):\n        x, y, z = center\n        slice_idx = x  # Visualize the xy-plane slice at x-coordinate\n\n        # Extract the slice from augmented data\n        slice_data_augmented = data[slice_idx]\n        # Extract the slice from original data\n        slice_data_original = original_data[slice_idx]\n\n        # Display the slice without highlight from augmented data\n        axes[i, 0].imshow(slice_data_augmented, cmap='gray')\n        axes[i, 0].set_title(f\"Center: {center} (Augmented Tomogram Slice)\")\n        axes[i, 0].axis('off')\n\n        # Display the slice with highlight from augmented data\n        axes[i, 1].imshow(slice_data_augmented, cmap='gray')\n        circle = plt.Circle((z, y), radius+3, color='red', fill=False, linewidth=1)\n        axes[i, 1].add_patch(circle)\n        axes[i, 1].set_title(f\"Center: {center} (Augmented Tomogram Slice)\")\n        axes[i, 1].axis('off')\n\n        # Display the slice without highlight from original data\n        axes[i, 2].imshow(slice_data_original, cmap='gray')\n        axes[i, 2].set_title(f\"Center: {center} (Original Tomogram)\")\n        axes[i, 2].axis('off')\n\n    plt.suptitle(title)\n    plt.tight_layout()\n    plt.show()\n\nexp_id = 1\nexperiment = experiment_mapping[exp_id]\ndata = tomograms[exp_id]['1'][:]\naugmented_data = augmented_tomograms[exp_id]\njson_file_path = os.path.join(overlay_base_path, experiment, \"Picks\", \"apo-ferritin.json\")\ncenters = load_centers(json_file_path, voxel_spacing)\nadjusted_centers = [(z // resolution, y // resolution, x // resolution) for z, y, x in centers]\naugmented_targets = final_targets[exp_id]\n\n# Visualize 2 original centers\nprint(\"**********. (Original Centers) *********\")\nvisualize_apo_ferritin(\n    data=augmented_data, \n    original_data=data, \n    centers=adjusted_centers[:2], \n    radius=voxel_radius, \n    title=\"\"\n)\n\nprint(\"**********. (New Centers) *********\")\n# Visualize 5 new centers from final_targets\nvisualize_apo_ferritin(\n    data=augmented_data, \n    original_data=data, \n    centers=[target for target in augmented_targets if target not in adjusted_centers][:5],\n    radius=voxel_radius, \n    title=\"\"\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:19.297155Z","iopub.execute_input":"2025-01-30T23:39:19.297394Z","iopub.status.idle":"2025-01-30T23:39:22.553051Z","shell.execute_reply.started":"2025-01-30T23:39:19.297373Z","shell.execute_reply":"2025-01-30T23:39:22.551836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tomograms = []\nfor exp_id in range(1, 8):\n    if exp_id < 6:\n        # Load augmented tomogram data\n        tomogram = augmented_tomograms[exp_id]\n    else:\n        tomogram = tomograms[exp_id]['1'][:]\n        \n    train_tomograms.append(tomogram)\n\ntarget_centers = []\nfor exp_id in range(1, 8):\n    if exp_id < 6:\n        adjusted_centers = final_targets[exp_id]\n    else:\n        json_file_path = os.path.join(overlay_base_path, experiment, \"Picks\", \"apo-ferritin.json\")\n        centers = load_centers(json_file_path, voxel_spacing)\n        adjusted_centers = [(z // resolution, y // resolution, x // resolution) for z, y, x in centers]\n        \n    target_centers.append(adjusted_centers)\n\n\nfor idx in range(0, 7):\n    print(type(train_tomograms[idx]), train_tomograms[idx].shape)\n    print(type(target_centers[idx]), len(target_centers[idx]))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:22.554495Z","iopub.execute_input":"2025-01-30T23:39:22.554914Z","iopub.status.idle":"2025-01-30T23:39:22.900063Z","shell.execute_reply.started":"2025-01-30T23:39:22.554870Z","shell.execute_reply":"2025-01-30T23:39:22.899241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split into training and validation datasets\ntrain_data = train_tomograms[:5]\ntrain_centers = target_centers[:5]\n\nval_data = train_tomograms[5:]\nval_centers = target_centers[5:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:22.900821Z","iopub.execute_input":"2025-01-30T23:39:22.901102Z","iopub.status.idle":"2025-01-30T23:39:22.904995Z","shell.execute_reply.started":"2025-01-30T23:39:22.901077Z","shell.execute_reply":"2025-01-30T23:39:22.904074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Create a binary mask for a 3D tomogram based on given centers and radius.\ndef create_binary_mask(tomogram_shape, centers, radius):\n    \n    mask = np.zeros(tomogram_shape, dtype=np.uint8)\n\n    for center in centers:\n    #for center in centers:\n        x, y, z = center\n        xx, yy, zz = np.ogrid[:tomogram_shape[0], :tomogram_shape[1], :tomogram_shape[2]]\n        distance = np.sqrt((xx - x)**2 + (yy - y)**2 + (zz - z)**2)\n        mask[distance <= radius] = 1\n    return mask\n\n# Generate binary masks for training data\ntrain_masks = []\nfor i, (data, centers) in enumerate(tqdm(zip(train_tomograms[:5], target_centers[:5]), desc='Generating Training Masks')):\n    mask = create_binary_mask(data.shape, centers, voxel_radius)\n    train_masks.append(mask)\n\n# Generate and save binary masks for validation data\nval_masks = []\nfor i, (data, centers) in enumerate(tqdm(zip(train_tomograms[5:], target_centers[5:]), desc='Generating Validation Masks')):\n    mask = create_binary_mask(data.shape, centers, voxel_radius)\n    val_masks.append(mask)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T23:39:22.905869Z","iopub.execute_input":"2025-01-30T23:39:22.906156Z","iopub.status.idle":"2025-01-31T00:16:11.958443Z","shell.execute_reply.started":"2025-01-30T23:39:22.906126Z","shell.execute_reply":"2025-01-31T00:16:11.957607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Normalize the data to a range [min, max].\ndef min_max_normalize(data, new_min=-1, new_max=1):\n    \n    min_val = data.min()\n    max_val = data.max()\n    \n    # Avoid division by zero\n    if max_val - min_val == 0:\n        return np.zeros_like(data)\n    \n    return (data - min_val) / (max_val - min_val) * (new_max - new_min) + new_min\n\n# Normalize train_data\nnormalized_train_data = [min_max_normalize(tomogram, -1, 1) for tomogram in train_data]\n# Normalize val_data\nnormalized_val_data = [min_max_normalize(tomogram, -1, 1) for tomogram in val_data]\n\n# Check normalization\nprint(f\"Train Data Range: {normalized_train_data[0].min()} to {normalized_train_data[0].max()}\")\nprint(f\"Validation Data Range: {normalized_val_data[0].min()} to {normalized_val_data[0].max()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-31T00:16:11.959351Z","iopub.execute_input":"2025-01-31T00:16:11.959678Z","iopub.status.idle":"2025-01-31T00:16:12.214214Z","shell.execute_reply.started":"2025-01-31T00:16:11.959644Z","shell.execute_reply":"2025-01-31T00:16:12.213351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Training\n\nimport tensorflow as tf\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'  # Suppress TensorFlow logging\ntf.random.set_seed(42)\nnp.random.seed(42)\n\n################################################################################\n# 2. DICE COEFFICIENT AND LOSS FUNCTIONS\n################################################################################\n#Computes the Dice coefficient between the true labels and predicted labels.\ndef dice_coefficient(y_true, y_pred, smooth=1e-6):\n    \n    y_true_flat = tf.reshape(y_true, [-1])\n    y_pred_flat = tf.reshape(y_pred, [-1])\n    intersection = tf.reduce_sum(y_true_flat * y_pred_flat)\n    coefficient = (2. * intersection + smooth) / (\n        tf.reduce_sum(y_true_flat) + tf.reduce_sum(y_pred_flat) + smooth\n    )\n    return coefficient\n\n#Dice loss for binary segmentation, minimizing 1 - dice_coefficient.\ndef dice_loss(y_true, y_pred):\n    \n    return 1.0 - dice_coefficient(y_true, y_pred)\n\ndef combined_bce_dice_loss(y_true, y_pred):\n    \"\"\"\n    Combines Binary Cross-Entropy with Dice loss to balance\n    voxel-level accuracy (BCE) with region-level overlap (Dice).\n    \"\"\"\n    bce = tf.keras.losses.binary_crossentropy(y_true, y_pred)\n    return bce + dice_loss(y_true, y_pred)\n\n################################################################################\n# 3. METRICS: PRECISION, RECALL, AND F1-SCORE\n################################################################################\nclass F1Score(tf.keras.metrics.Metric):\n    \"\"\"\n    Custom F1 Score metric using the Precision and Recall from Keras.\n    \"\"\"\n    def __init__(self, name='f1_score', **kwargs):\n        super(F1Score, self).__init__(name=name, **kwargs)\n        self.precision = tf.keras.metrics.Precision()\n        self.recall = tf.keras.metrics.Recall()\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        y_pred_thresholded = tf.cast(y_pred > 0.5, tf.int32)\n        self.precision.update_state(y_true, y_pred_thresholded)\n        self.recall.update_state(y_true, y_pred_thresholded)\n\n    def result(self):\n        p = self.precision.result()\n        r = self.recall.result()\n        return 2 * p * r / (p + r + 1e-6)\n\n    def reset_states(self):\n        self.precision.reset_states()\n        self.recall.reset_states()\n\n################################################################################\n# 4. MODEL ARCHITECTURE (3D U-NET STYLE)\n################################################################################\n\ndef create_3d_unet_with_padding(input_shape=(92, 315, 315, 1)):\n    inputs = tf.keras.Input(shape=input_shape)\n\n    # -- Encoder --\n    c1 = tf.keras.layers.Conv3D(16, 3, activation='relu', padding='same')(inputs)\n    c1 = tf.keras.layers.Conv3D(16, 3, activation='relu', padding='same')(c1)\n    p1 = tf.keras.layers.MaxPooling3D((2,2,2))(c1)\n    # p1 shape: (None, 46, 157, 157, 16)\n\n    c2 = tf.keras.layers.Conv3D(32, 3, activation='relu', padding='same')(p1)\n    c2 = tf.keras.layers.Conv3D(32, 3, activation='relu', padding='same')(c2)\n    p2 = tf.keras.layers.MaxPooling3D((2,2,2))(c2)\n    # p2 shape: (None, 23, 78, 78, 32)\n\n    c3 = tf.keras.layers.Conv3D(64, 3, activation='relu', padding='same')(p2)\n    c3 = tf.keras.layers.Conv3D(64, 3, activation='relu', padding='same')(c3)\n    # c3 shape: (None, 23, 78, 78, 64)\n\n    # -- Decoder --\n    # Up-sample from c3\n    u2 = tf.keras.layers.UpSampling3D((2,2,2))(c3)\n    # shape: (None, 46, 156, 156, 64)\n\n    # Zero-pad by 1 voxel on width & height to go 156 -> 157\n    u2 = tf.keras.layers.ZeroPadding3D(padding=((0,0), (0,1), (0,1)))(u2)\n    # shape now: (None, 46, 157, 157, 64)\n\n    concat_2 = tf.keras.layers.Concatenate()([u2, c2])\n    c4 = tf.keras.layers.Conv3D(32, 3, activation='relu', padding='same')(concat_2)\n    c4 = tf.keras.layers.Conv3D(32, 3, activation='relu', padding='same')(c4)\n    # c4 shape: (None, 46, 157, 157, 32)\n\n    # Up-sample from c4\n    u1 = tf.keras.layers.UpSampling3D((2,2,2))(c4)\n    # shape: (None, 92, 314, 314, 32)\n\n    # Zero-pad by 1 voxel on width & height to go 314 -> 315\n    u1 = tf.keras.layers.ZeroPadding3D(padding=((0,0), (0,1), (0,1)))(u1)\n    # shape now: (None, 92, 315, 315, 32)\n\n    concat_1 = tf.keras.layers.Concatenate()([u1, c1])\n    c5 = tf.keras.layers.Conv3D(16, 3, activation='relu', padding='same')(concat_1)\n    c5 = tf.keras.layers.Conv3D(16, 3, activation='relu', padding='same')(c5)\n\n    outputs = tf.keras.layers.Conv3D(1, 1, activation='sigmoid')(c5)\n    model = tf.keras.Model(inputs=[inputs], outputs=[outputs])\n    return model\n\n\nmodel = create_3d_unet_with_padding()\n#model.summary()\n\n################################################################################\n# 5. COMPILE THE MODEL WITH OUR LOSS AND METRICS\n################################################################################\n\n# Define custom metrics\nprecision = tf.keras.metrics.Precision(name='precision')\nrecall = tf.keras.metrics.Recall(name='recall')\nf1_score = F1Score(name='f1_score')\n\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    loss=combined_bce_dice_loss,\n    metrics=[precision, recall, f1_score]\n)\n\n################################################################################\n# 6. CREATE TFDATA DATASETS FOR TRAINING AND VALIDATION\n################################################################################\n\ndef list_to_5d_array(vol_list, as_float32=False):\n    \"\"\"\n    vol_list: list of NumPy arrays (Z, H, W).\n    If as_float32=True, convert each to float32.\n    \"\"\"\n    import numpy as np\n    vols_4d = []\n    for vol in vol_list:\n        if as_float32:\n            vol = vol.astype(np.float32)\n        vol_4d = np.expand_dims(vol, axis=-1)  # => (Z, H, W, 1)\n        vols_4d.append(vol_4d)\n    # Stack => (N, Z, H, W, 1)\n    return np.stack(vols_4d, axis=0)\n\n#add RGB channel (gray) as last dimension \ntrain_data_5d = list_to_5d_array(normalized_train_data, as_float32=True)     # shape: (5, 92, 315, 315, 1)\ntrain_masks_5d = list_to_5d_array(train_masks, as_float32=True)   # shape: (5, 92, 315, 315, 1)\n\nval_data_5d = list_to_5d_array(normalized_val_data, as_float32=True)     # shape: (2, 92, 315, 315, 1)\nval_masks_5d = list_to_5d_array(val_masks, as_float32=True)   # shape: (2, 92, 315, 315, 1)\n\n\nprint(\"train_data_5d shape:\", train_data_5d.shape)  \nprint(\"train_masks_5d shape:\", train_masks_5d.shape)\nprint(\"val_data_5d shape:\", val_data_5d.shape)\nprint(\"val_masks_5d shape:\", val_masks_5d.shape)\n\n\n################################################################################\n# 7. TRAINING THE MODEL\n################################################################################\n\nBATCH_SIZE = 1  # Typically 1 for full-volume 3D\nEPOCHS = 10\n\ntrain_dataset = tf.data.Dataset.from_tensor_slices((train_data_5d, train_masks_5d))\ntrain_dataset = train_dataset.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)\n\nval_dataset = tf.data.Dataset.from_tensor_slices((val_data_5d, val_masks_5d))\nval_dataset = val_dataset.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)\n\nhistory = model.fit(\n    train_dataset,\n    epochs=EPOCHS,\n    validation_data=val_dataset\n)\n\n\nprint(\"\\nTraining completed.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-31T00:16:12.215330Z","iopub.execute_input":"2025-01-31T00:16:12.215633Z","iopub.status.idle":"2025-01-31T00:22:24.112812Z","shell.execute_reply.started":"2025-01-31T00:16:12.215609Z","shell.execute_reply":"2025-01-31T00:22:24.111997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"################################################################################\n# VISUALIZATION OF METRICS AND PREDICTIONS\n################################################################################\n\nplt.figure(figsize=(12, 4))\n\n# Loss\nplt.subplot(1, 3, 1)\nplt.plot(history.history['loss'], label='Train Loss')\nplt.plot(history.history['val_loss'], label='Val Loss')\nplt.title(\"Loss\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\n# Precision\nplt.subplot(1, 3, 2)\nplt.plot(history.history['precision'], label='Train Precision')\nplt.plot(history.history['val_precision'], label='Val Precision')\nplt.title(\"Precision\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Precision\")\nplt.legend()\n\n# Recall\nplt.subplot(1, 3, 3)\nplt.plot(history.history['recall'], label='Train Recall')\nplt.plot(history.history['val_recall'], label='Val Recall')\nplt.title(\"Recall\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Recall\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n\n# ------------------------------------------------------------------------------\n# Convert val_data (list) into 5D NumPy array, then visualize predictions\n# ------------------------------------------------------------------------------\n\n# val_data is a list of shape N, each element (92, 315, 315)\n# val_masks is a list of shape N, each element (92, 315, 315)\n\n# 1) Expand dims on each volume to add a channel => (92, 315, 315, 1)\nval_data_4d = [np.expand_dims(vol, axis=-1) for vol in normalized_val_data]\nval_masks_4d = [np.expand_dims(mask, axis=-1) for mask in val_masks]\n\n# 2) Stack into a single array => (N, 92, 315, 315, 1)\nval_data_5d = np.stack(val_data_4d, axis=0)\nval_masks_5d = np.stack(val_masks_4d, axis=0)\n\nprint(\"val_data_5d shape:\", val_data_5d.shape)   # Expect (N, 92, 315, 315, 1)\nprint(\"val_masks_5d shape:\", val_masks_5d.shape) # Expect (N, 92, 315, 315, 1)\n\n# Take a single volume from val_data_5d\nsample_idx = 0\nsample_vol = val_data_5d[sample_idx:sample_idx+1]  # shape: (1, 92, 315, 315, 1)\ntrue_mask = val_masks_5d[sample_idx]               # shape: (92, 315, 315, 1)\n\n# Get model predictions => (1, 92, 315, 315, 1)\npred_mask_prob = model.predict(sample_vol)\npred_mask = (pred_mask_prob > 0.5).astype(np.uint8)[0]  # => (92, 315, 315, 1)\n\n# Display a 2D slice (example: z=46)\nz_slice = 35\n\nplt.figure(figsize=(16, 5))\n\n# Tomogram slice\nplt.subplot(1, 3, 1)\nplt.imshow(sample_vol[0, z_slice, :, :, 0], cmap='gray')\nplt.title(\"Validation Tomogram Slice\")\n\n# Ground-truth mask slice\nplt.subplot(1, 3, 2)\nplt.imshow(true_mask[z_slice, :, :, 0], cmap='gray')\nplt.title(\"Ground Truth Mask\")\n\n# Predicted mask slice\nplt.subplot(1, 3, 3)\nplt.imshow(pred_mask[z_slice, :, :, 0], cmap='gray')\nplt.title(\"Predicted Mask\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-31T00:35:59.425643Z","iopub.execute_input":"2025-01-31T00:35:59.425937Z","iopub.status.idle":"2025-01-31T00:36:01.640290Z","shell.execute_reply.started":"2025-01-31T00:35:59.425915Z","shell.execute_reply":"2025-01-31T00:36:01.639405Z"}},"outputs":[],"execution_count":null}]}