{"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":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":229899399,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Comparison of Denoising Methods for Cryo-ET Data\n\nThis notebook compares various denoising methods for bacterial flagellar motor data from cryo-electron tomography (cryo-ET). We've added Noise2Void (N2V) to the comparison alongside traditional methods like BM3D, OpenCV NLMeans, Gaussian blur, Wiener filter, and a crude 3D approach.\n\n### Key Observations on Noise2Void Performance:\n\n1. **Processing Speed**: N2V is approximately 13.5× faster than BM3D (73 seconds vs. 986 seconds for 30 slices) while achieving the highest noise reduction at 49.3%.\n\n2. **Noise Reduction Effectiveness**: N2V consistently shows the best noise reduction percentages across all tested images:\n   - Achieving up to 84% noise reduction in some samples\n   - Significantly outperforming Gaussian blur (43.1%), BM3D (22.6%), and Wiener filtering (20.2%)\n\n3. **Detail Preservation vs. Smoothing Trade-offs**:\n   - N2V excels at preserving bacterial cell structures and flagellar motors while dramatically reducing background noise\n   - In some cases (particularly with continuous linear structures), N2V can over-smooth fine details\n   - BM3D often provides better preservation of continuous linear features but with less noise reduction\n\n4. **Image-specific Performance**:\n   - N2V performs exceptionally well on \"bean-shaped\" flagellar motor structures \n   - BM3D may preserve certain fine details better in highly textured regions\n   - Image-specific training of N2V can further improve performance on particular structural features\n\n5. **Optimization Potential**: Training N2V on small batches (3-5 images) for fewer epochs (4-6) provides an excellent balance between:\n   - Training time (approximately 12-30 minutes per batch)\n   - Inference time (approximately 2.4 seconds per slice, 10 minutes for 300 slices)\n   - Quality (customized to specific structural features)\n\n### Practical Implementation Considerations:\n\n- For exploratory analysis and initial processing, N2V provides the best balance of speed and quality\n- For publication-quality final images of critical structures, BM3D might still be preferred for certain detail types\n- A hybrid approach could be ideal - using N2V for initial processing and BM3D selectively on regions with fine linear structures\n- Custom training of N2V on specific parts of the tomogram further improves results and addresses detail preservation issues\n\nOur current optimized workflow uses small-batch, few-epoch training of N2V, which produces superior results while being 7× faster than BM3D, making it an excellent choice for processing large cryo-ET datasets.","metadata":{}},{"cell_type":"code","source":"# Install required packages\nfrom IPython.display import clear_output\n!pip install -q bm3d n2v tensorflow==2.13.0 csbdeep pillow tqdm\nclear_output()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport glob\nimport time\nimport cv2\nfrom skimage import img_as_float\nimport bm3d\nfrom scipy import signal\nfrom n2v.models import N2V\nimport warnings\nfrom tqdm.notebook import tqdm\n\n# Suppress warnings and tqdm output\nwarnings.filterwarnings('ignore')\ntqdm.__init__ = lambda *args, **kwargs: None\ntqdm.update = lambda *args, **kwargs: None\ntqdm.close = lambda *args, **kwargs: None\ntqdm.__iter__ = lambda self: iter([])\n\n# Define paths\ndata_path = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_00e047'\nMODEL_PATH = '/kaggle/input/byu-denoising-cryo-et-with-noise2void/models/'\n\ndef load_volume(folder_path, max_slices=30, start_slice=0):\n    \"\"\"Load a subset of slices from the volume\"\"\"\n    all_files = sorted(glob.glob(os.path.join(folder_path, '*.jpg')))\n    selected_files = all_files[start_slice:start_slice+max_slices]\n    \n    volume = []\n    for file in selected_files:\n        img = np.array(Image.open(file).convert('L'))  # Convert to grayscale\n        volume.append(img)\n    \n    return np.array(volume)\n\ndef process_cv2_denoise(volume, h=10, template_window_size=7, search_window_size=21):\n    \"\"\"Process each slice with OpenCV's fastNlMeansDenoising\"\"\"\n    output = np.zeros_like(volume)\n    start_time = time.time()\n    \n    for i, slice_img in enumerate(volume):\n        output[i] = cv2.fastNlMeansDenoising(\n            slice_img, \n            None, \n            h=h,\n            templateWindowSize=template_window_size,\n            searchWindowSize=search_window_size\n        )\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef process_gaussian(volume, kernel_size=5):\n    \"\"\"Process each slice with Gaussian blur\"\"\"\n    output = np.zeros_like(volume)\n    start_time = time.time()\n    \n    for i, slice_img in enumerate(volume):\n        output[i] = cv2.GaussianBlur(slice_img, (kernel_size, kernel_size), 0)\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef process_bm3d(volume, sigma_psd=0.1):\n    \"\"\"Process volume with BM3D denoising\"\"\"\n    output = np.zeros_like(volume, dtype=np.float32)\n    start_time = time.time()\n    \n    for i, slice_img in enumerate(volume):\n        # Convert to float and normalize\n        slice_float = img_as_float(slice_img)\n        \n        # Apply BM3D denoising\n        denoised_slice = bm3d.bm3d(slice_float, sigma_psd=sigma_psd)\n        \n        # Convert back to uint8\n        output[i] = (denoised_slice * 255).astype(np.uint8)\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef process_wiener(volume, kernel_size=5, noise_power=0.01):\n    \"\"\"Process volume with Wiener filter\"\"\"\n    output = np.zeros_like(volume, dtype=np.float32)\n    start_time = time.time()\n    \n    for i, slice_img in enumerate(volume):\n        # Normalize to [0, 1] range\n        slice_norm = slice_img.astype(np.float32) / 255.0\n        \n        # Create a local mean filter\n        kernel = np.ones((kernel_size, kernel_size)) / (kernel_size**2)\n        \n        # Compute local mean\n        img_mean = signal.convolve2d(slice_norm, kernel, mode='same')\n        \n        # Compute local variance\n        img_sqr_mean = signal.convolve2d(slice_norm**2, kernel, mode='same')\n        img_var = img_sqr_mean - img_mean**2\n        \n        # Ensure variance is positive\n        img_var = np.maximum(img_var, 0)\n        \n        # Apply Wiener filter formula\n        denoised = img_mean + ((img_var - noise_power) / np.maximum(img_var, noise_power)) * (slice_norm - img_mean)\n        \n        # Clip values to valid range and convert back to 8-bit\n        denoised = np.clip(denoised, 0, 1)\n        output[i] = (denoised * 255).astype(np.uint8)\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef process_crude_3d(volume, kernel_size=3):\n    \"\"\"Process with crude 3D denoising by averaging neighboring slices\"\"\"\n    output = np.zeros_like(volume)\n    start_time = time.time()\n    \n    for i in range(volume.shape[0]):\n        # Get neighboring slices\n        slice_min = max(0, i - kernel_size//2)\n        slice_max = min(volume.shape[0], i + kernel_size//2 + 1)\n        # Average the neighboring slices\n        neighbors = volume[slice_min:slice_max]\n        # Apply 2D denoising to each slice and then average\n        denoised_neighbors = []\n        for neighbor in neighbors:\n            denoised = cv2.fastNlMeansDenoising(neighbor, None, h=10)\n            denoised_neighbors.append(denoised)\n        output[i] = np.mean(denoised_neighbors, axis=0).astype(np.uint8)\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef process_noise2void(volume):\n    \"\"\"Process volume with pre-trained Noise2Void model\"\"\"\n    output = np.zeros_like(volume, dtype=np.uint8)\n    start_time = time.time()\n    \n    # Redirect standard output to suppress N2V loading/processing messages\n    import os, sys\n    original_stdout = sys.stdout\n    sys.stdout = open(os.devnull, 'w')\n    \n    try:\n        # Load the pre-trained Noise2Void model\n        model = N2V(config=None, name='n2v_cryoET_8slices', basedir=MODEL_PATH)\n        \n        for i, slice_img in enumerate(volume):\n            # Add dimensions for N2V (SYXC format where S is sample/batch dimension)\n            img_for_pred = slice_img[np.newaxis, ..., np.newaxis]\n            \n            # Apply Noise2Void denoising\n            denoised = model.predict(img_for_pred, axes='SYXC')\n            \n            # Remove the batch and channel dimensions\n            denoised_img = denoised[0, ..., 0]\n            \n            # Convert to uint8 for display\n            if denoised_img.dtype != np.uint8:\n                denoised_img = np.clip(denoised_img, 0, 255).astype(np.uint8)\n                \n            output[i] = denoised_img\n    finally:\n        # Restore standard output\n        sys.stdout.close()\n        sys.stdout = original_stdout\n    \n    total_time = time.time() - start_time\n    return output, total_time\n\ndef display_results(original, results_dict, slice_idx=None):\n    \"\"\"Display comparison of original and processed results\"\"\"\n    if slice_idx is None:\n        slice_idx = original.shape[0] // 2  # Middle slice\n    \n    num_results = len(results_dict) + 1  # +1 for original\n    fig_width = 6 * num_results\n    \n    plt.figure(figsize=(fig_width, 6))\n    \n    # Display original\n    plt.subplot(1, num_results, 1)\n    plt.imshow(original[slice_idx], cmap='gray')\n    plt.title('Original')\n    plt.axis('off')\n    \n    # Display results\n    for i, (name, result) in enumerate(results_dict.items(), 2):\n        plt.subplot(1, num_results, i)\n        plt.imshow(result[slice_idx], cmap='gray')\n        plt.title(name)\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.savefig('denoising_comparison.png')\n    plt.show()\n\ndef run_denoising_comparison(data_path, max_slices=30, start_slice=120):\n    \"\"\"Run full denoising comparison with minimal output\"\"\"\n    # Load the volume data quietly\n    volume = load_volume(data_path, max_slices=max_slices, start_slice=start_slice)\n    \n    # Store results and times for comparison\n    results = {}\n    times = {}\n    \n    # Suppress output during processing\n    import sys\n    original_stdout = sys.stdout\n    null_output = open(os.devnull, 'w')\n    sys.stdout = null_output\n    \n    try:\n        # Run the different denoising methods\n        results['OpenCV NLMeans'], times['OpenCV NLMeans'] = process_cv2_denoise(volume)\n        results['Gaussian'], times['Gaussian'] = process_gaussian(volume)\n        results['BM3D'], times['BM3D'] = process_bm3d(volume)\n        results['Wiener'], times['Wiener'] = process_wiener(volume)\n        results['Crude 3D'], times['Crude 3D'] = process_crude_3d(volume)\n        results['Noise2Void'], times['Noise2Void'] = process_noise2void(volume)\n    finally:\n        # Restore stdout\n        sys.stdout = original_stdout\n        null_output.close()\n    \n    # Display the results\n    display_results(volume, results)\n    \n    # Calculate noise reduction statistics\n    slice_idx = volume.shape[0] // 2  # Middle slice\n    original_noise = np.std(volume[slice_idx])\n    noise_reduction = {}\n    \n    for name, result in results.items():\n        noise_after = np.std(result[slice_idx])\n        reduction_percent = (1 - noise_after/original_noise) * 100\n        noise_reduction[name] = reduction_percent\n    \n    # Print just the key information\n    print(f\"Volume shape: {volume.shape}\")\n    print(\"\\nProcessing Time Comparison:\")\n    for name, time_value in times.items():\n        print(f\"{name}: {time_value:.2f} seconds\")\n    \n    print(\"\\nNoise Reduction Statistics:\")\n    for name, reduction in noise_reduction.items():\n        print(f\"{name}: {reduction:.1f}% noise reduction\")\n    \n    # Create a zoomed crop view for clearer comparison\n    center_slice = volume.shape[0] // 2\n    h, w = volume[center_slice].shape\n    center_y, center_x = h//2, w//2\n    crop_size = 200\n    \n    # Create a new figure for the cropped view\n    plt.figure(figsize=(20, 10))\n    \n    # Original cropped\n    plt.subplot(1, len(results) + 1, 1)\n    crop_original = volume[center_slice][center_y-crop_size//2:center_y+crop_size//2, \n                         center_x-crop_size//2:center_x+crop_size//2]\n    plt.imshow(crop_original, cmap='gray')\n    plt.title('Original (Center Crop)')\n    plt.axis('off')\n    \n    # Processed cropped\n    for i, (name, result) in enumerate(results.items(), 2):\n        crop_result = result[center_slice][center_y-crop_size//2:center_y+crop_size//2, \n                            center_x-crop_size//2:center_x+crop_size//2]\n        plt.subplot(1, len(results) + 1, i)\n        plt.imshow(crop_result, cmap='gray')\n        plt.title(f'{name} (Center Crop)')\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.savefig('denoising_comparison_cropped.png')\n    plt.show()\n    \n    return volume, results, times, noise_reduction\n\n\nvolume, results, times, noise_reduction = run_denoising_comparison(data_path)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-31T19:34:56.123274Z","iopub.execute_input":"2025-03-31T19:34:56.123578Z","iopub.status.idle":"2025-03-31T19:52:07.607499Z","shell.execute_reply.started":"2025-03-31T19:34:56.123542Z","shell.execute_reply":"2025-03-31T19:52:07.606285Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Wiener Filter\nThe Wiener filter represents an adaptive denoising approach particularly well-suited for Cryo-ET data. Unlike simple smoothing techniques, this method optimally balances noise reduction and feature preservation by adapting to local image statistics. The filter estimates the signal-to-noise ratio across different regions, applying stronger filtering in homogeneous areas while preserving edges and fine structures critical for flagellar motor identification. Our implementation allows for parameter tuning to accommodate varying noise levels in tomographic slices, making it computationally efficient for large Cryo-ET datasets while maintaining excellent structural detail.","metadata":{}},{"cell_type":"code","source":"def wiener_filter(img, kernel_size=5, noise_power=0.012):\n    \"\"\"\n    Apply Wiener filter for denoising a grayscale image.\n    \n    Parameters:\n        img: Input grayscale image\n        kernel_size: Size of the kernel (odd number)\n        noise_power: Estimated noise power\n        \n    Returns:\n        Denoised image\n    \"\"\"\n    img_norm = img.astype(np.float32) / 255.0\n    \n    # Create a local mean filter\n    kernel = np.ones((kernel_size, kernel_size)) / (kernel_size**2)\n    \n    # Compute local mean\n    img_mean = signal.convolve2d(img_norm, kernel, mode='same')\n    \n    # Compute local variance\n    img_sqr_mean = signal.convolve2d(img_norm**2, kernel, mode='same')\n    img_var = img_sqr_mean - img_mean**2\n    \n    # Ensure variance is positive\n    img_var = np.maximum(img_var, 0)\n    \n    # Apply Wiener filter formula\n    denoised = img_mean + ((img_var - noise_power) / np.maximum(img_var, noise_power)) * (img_norm - img_mean)\n    \n    denoised = np.clip(denoised, 0, 1)\n    \n    denoised_img = (denoised * 255).astype(np.uint8)\n    \n    return denoised_img\n\nto_show={'00e047':169,\n        '00e463':222,\n        '1da097':34}\n\nfor tomo_id, Motor_axis_0 in to_show.items():\n    input_image = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_' + tomo_id + '/slice_' + str(Motor_axis_0).zfill(4) + '.jpg'\n    img = cv2.imread(input_image, cv2.IMREAD_GRAYSCALE)\n    \n    denoised_img = wiener_filter(img, kernel_size=11, noise_power=0.02)\n    \n    # Display full images results\n    fig, axes = plt.subplots(1, 2, figsize=(15, 8))\n    axes[0].imshow(img, cmap='gray')\n    axes[0].set_title(f'Noisy Image (tomo_{tomo_id}, slice_{Motor_axis_0})')\n    axes[0].axis('off')\n    \n    axes[1].imshow(denoised_img, cmap='gray')\n    axes[1].set_title('Denoised Image (Wiener Filter)')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Create and display center crops for detailed comparison\n    h, w = img.shape\n    center_y, center_x = h//2, w//2\n    crop_size = 200\n    \n    # Extract center crops\n    crop_original = img[center_y-crop_size//2:center_y+crop_size//2, \n                        center_x-crop_size//2:center_x+crop_size//2]\n    crop_denoised = denoised_img[center_y-crop_size//2:center_y+crop_size//2, \n                               center_x-crop_size//2:center_x+crop_size//2]\n    \n    # Display crops\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    axes[0].imshow(crop_original, cmap='gray')\n    axes[0].set_title('Original (Center Crop)')\n    axes[0].axis('off')\n    \n    axes[1].imshow(crop_denoised, cmap='gray')\n    axes[1].set_title('Denoised (Center Crop)')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T20:23:53.656069Z","iopub.execute_input":"2025-03-31T20:23:53.656511Z","iopub.status.idle":"2025-03-31T20:24:00.155261Z","shell.execute_reply.started":"2025-03-31T20:23:53.656483Z","shell.execute_reply":"2025-03-31T20:24:00.154039Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## BM3D (Block-Matching and 3D filtering)\nBM3D stands as the state-of-the-art denoising algorithm in our comparative analysis, delivering superior quality for Cryo-ET applications. This sophisticated method works by grouping similar 2D image patches into 3D arrays, applying collaborative filtering in a transform domain, and then reconstructing the denoised result. For bacterial flagellar motor detection, BM3D excels at preserving the fine structural details while significantly reducing noise, enhancing downstream YOLO model performance. Though computationally more intensive than other methods, the quality improvement justifies its use for datasets where maximum detail preservation is required.","metadata":{}},{"cell_type":"code","source":"from skimage import io, img_as_float, img_as_ubyte\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport bm3d\n\nto_show={'00e047':169,\n        '00e463':222,\n        '1da097':34}\n\nfor tomo_id,Motor_axis_0 in to_show.items():\n    input_image = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_' + tomo_id + '/slice_' + str(Motor_axis_0).zfill(4) + '.jpg'\n    \n    noisy_image = img_as_float(io.imread(input_image, as_gray=True))\n    denoised_image = bm3d.bm3d(noisy_image, sigma_psd=0.15)\n    \n    # Display results\n    fig, axes = plt.subplots(1, 2, figsize=(15, 8))\n    axes[0].imshow(noisy_image, cmap='gray')\n    axes[0].set_title('Noisy Image')\n    axes[0].axis('off')\n    \n    axes[1].imshow(denoised_image, cmap='gray')\n    axes[1].set_title('Denoised Image (BM3D)')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Show a zoomed-in crop for better comparison\n    h, w = noisy_image.shape\n    center_y, center_x = h//2, w//2\n    crop_size = 200\n    \n    # Extract center crops\n    crop_original = noisy_image[center_y-crop_size//2:center_y+crop_size//2, \n                              center_x-crop_size//2:center_x+crop_size//2]\n    crop_denoised = denoised_image[center_y-crop_size//2:center_y+crop_size//2, \n                                 center_x-crop_size//2:center_x+crop_size//2]\n    \n    # Display crops\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    axes[0].imshow(crop_original, cmap='gray')\n    axes[0].set_title('Original (Center Crop)')\n    axes[0].axis('off')\n    \n    axes[1].imshow(crop_denoised, cmap='gray')\n    axes[1].set_title('Denoised (Center Crop)')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T20:18:25.073985Z","iopub.execute_input":"2025-03-31T20:18:25.074524Z","iopub.status.idle":"2025-03-31T20:20:05.533601Z","shell.execute_reply.started":"2025-03-31T20:18:25.074492Z","shell.execute_reply":"2025-03-31T20:20:05.532249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Noise2Void Inference Section","metadata":{}},{"cell_type":"code","source":"!pip install -q n2v tensorflow==2.13.0 csbdeep pillow tqdm\n# Import necessary libraries\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport os\nfrom n2v.models import N2V\nfrom skimage import io\n\n# Define the paths to pre-trained model weights and config\nMODEL_PATH = '/kaggle/input/byu-denoising-cryo-et-with-noise2void/models/'\n\n# Load the pre-trained Noise2Void model\nmodel = N2V(config=None, name='n2v_cryoET_8slices', basedir=MODEL_PATH)\n\n# List of tomogram IDs and specific slices to process\nto_show = {\n    '00e047': 169,\n    '00e463': 222,\n    '1da097': 34\n}\n\n# Process and display each example\nfor tomo_id, motor_axis_0 in to_show.items():\n    # Construct input image path\n    input_image = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_' + tomo_id + '/slice_' + str(motor_axis_0).zfill(4) + '.jpg'\n    \n    # Load image\n    img = np.array(Image.open(input_image).convert('L'))\n    \n    # Add dimensions for N2V (SYXC format where S is sample/batch dimension)\n    img_for_pred = img[np.newaxis, ..., np.newaxis]\n    \n    # Apply Noise2Void denoising\n    denoised = model.predict(img_for_pred, axes='SYXC')\n    \n    # Remove the batch and channel dimensions\n    denoised_img = denoised[0, ..., 0]\n    \n    # Convert to uint8 for display\n    if denoised_img.dtype != np.uint8:\n        denoised_img = np.clip(denoised_img, 0, 255).astype(np.uint8)\n    \n    # Calculate noise reduction statistics\n    noise_before = np.std(img)\n    noise_after = np.std(denoised_img)\n    reduction_percent = (1 - noise_after/noise_before) * 100\n    \n    # Display results\n    fig, axes = plt.subplots(1, 2, figsize=(15, 8))\n    axes[0].imshow(img, cmap='gray')\n    axes[0].set_title(f'Noisy Image (tomo_{tomo_id}, slice_{motor_axis_0})')\n    axes[0].axis('off')\n    \n    axes[1].imshow(denoised_img, cmap='gray')\n    axes[1].set_title(f'Denoised Image (Noise2Void) - {reduction_percent:.1f}% noise reduction')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Optional: Show a zoomed-in region for better comparison\n    # Select center region (200x200 pixels)\n    h, w = img.shape\n    center_y, center_x = h//2, w//2\n    crop_size = 200\n    \n    crop_original = img[center_y-crop_size//2:center_y+crop_size//2, \n                        center_x-crop_size//2:center_x+crop_size//2]\n    crop_denoised = denoised_img[center_y-crop_size//2:center_y+crop_size//2, \n                               center_x-crop_size//2:center_x+crop_size//2]\n    \n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    axes[0].imshow(crop_original, cmap='gray')\n    axes[0].set_title('Original (Center Crop)')\n    axes[0].axis('off')\n    \n    axes[1].imshow(crop_denoised, cmap='gray')\n    axes[1].set_title('Denoised (Center Crop)')\n    axes[1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n# Batch processing function for denoising multiple slices\ndef batch_process_n2v(input_dir, output_dir, slice_range=None):\n    \"\"\"\n    Process a batch of slices with the pre-trained Noise2Void model\n    \n    Parameters:\n        input_dir: Directory containing input JPEG slices\n        output_dir: Directory to save denoised slices\n        slice_range: Optional tuple (start, end) to process only a subset of slices\n    \"\"\"\n    # Create output directory if it doesn't exist\n    os.makedirs(output_dir, exist_ok=True)\n    \n    # Get all JPEG files in the input directory\n    all_files = sorted([f for f in os.listdir(input_dir) if f.endswith('.jpg')])\n    \n    # Apply slice range if specified\n    if slice_range is not None:\n        start, end = slice_range\n        all_files = all_files[start:end]\n    \n    print(f\"Processing {len(all_files)} slices...\")\n    \n    # Process each slice\n    for filename in all_files:\n        input_path = os.path.join(input_dir, filename)\n        output_path = os.path.join(output_dir, f\"denoised_{filename}\")\n        \n        # Load image\n        img = np.array(Image.open(input_path).convert('L'))\n        \n        # Add dimensions for N2V (SYXC format)\n        img_for_pred = img[np.newaxis, ..., np.newaxis]\n        \n        # Apply Noise2Void denoising\n        denoised = model.predict(img_for_pred, axes='SYXC')\n        \n        # Remove batch and channel dimensions\n        denoised_img = denoised[0, ..., 0]\n        \n        # Convert to uint8 for saving\n        if denoised_img.dtype != np.uint8:\n            denoised_img = np.clip(denoised_img, 0, 255).astype(np.uint8)\n        \n        # Save denoised image\n        Image.fromarray(denoised_img).save(output_path)\n    \n    print(f\"Batch processing complete. Denoised images saved to {output_dir}\")\n\n# Example usage of batch processing (commented out)\n# tomo_id = '00e047'\n# input_dir = f'/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_{tomo_id}'\n# output_dir = f'/kaggle/working/denoised_tomo_{tomo_id}'\n# batch_process_n2v(input_dir, output_dir, slice_range=(100, 200))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T20:15:21.540430Z","iopub.execute_input":"2025-03-31T20:15:21.540836Z","iopub.status.idle":"2025-03-31T20:15:38.841044Z","shell.execute_reply.started":"2025-03-31T20:15:21.540804Z","shell.execute_reply":"2025-03-31T20:15:38.839559Z"}},"outputs":[],"execution_count":null}]}