{"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,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Noise2Void Training on 8 Cryo-ET Slices (15 Epochs)\n# This script trains on 8 images and denoises samples from the full dataset\n\n# 1. Install required packages\n!pip install -q n2v tensorflow==2.13.0 csbdeep pillow tqdm\n\n# 1. Import dependencies\nimport os\nimport glob\nimport random\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom n2v.models import N2VConfig, N2V\nfrom n2v.internals.N2V_DataGenerator import N2V_DataGenerator\nfrom csbdeep.utils import plot_history\n\n# 2. Set up paths and parameters\nTOMO_PATH = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_00e047'\nOUTPUT_PATH = '/kaggle/working'\nMODEL_PATH = '/kaggle/working/models'\n\n# Create output directory if it doesn't exist\nos.makedirs(OUTPUT_PATH, exist_ok=True)\nos.makedirs(MODEL_PATH, exist_ok=True)\n\n# 3. Load a specific range of slices with a step\ndef load_slices(path, start_idx=0, end_idx=300, step=20):\n    \"\"\"Load a range of JPEG images from the path with specified indices and step\"\"\"\n    all_files = sorted(glob.glob(os.path.join(path, '*.jpg')))\n    \n    if len(all_files) == 0:\n        raise ValueError(f\"No JPEG files found in {path}\")\n    \n    print(f\"Found {len(all_files)} JPEG files in the directory\")\n    \n    # Make sure our range is valid\n    start_idx = max(0, min(start_idx, len(all_files) - 1))\n    end_idx = max(start_idx + 1, min(end_idx, len(all_files)))\n    \n    # Select files with the specified step\n    selected_indices = list(range(start_idx, end_idx, step))\n    selected_files = [all_files[i] for i in selected_indices]\n    \n    print(f\"Selected {len(selected_files)} slices from index {start_idx} to {end_idx-1} with step {step}\")\n    print(f\"Selected indices: {selected_indices}\")\n    \n    # Load the images into a list\n    images = []\n    filenames = []\n    \n    for file in tqdm(selected_files, desc=\"Loading slices\"):\n        # Load and convert to grayscale\n        img = np.array(Image.open(file).convert('L'))\n        images.append(img)\n        filenames.append(os.path.basename(file))\n    \n    return images, filenames, selected_files, all_files\n\n# Load slices with step 20 for testing\nimages, filenames, image_files, all_files = load_slices(TOMO_PATH, start_idx=0, end_idx=300, step=20)\n\n# 4. Display a sample of the loaded images\ndef display_sample_slices(images, filenames, num_samples=5):\n    # Select random samples if we have more than requested\n    if len(images) > num_samples:\n        indices = sorted(random.sample(range(len(images)), num_samples))\n    else:\n        indices = range(len(images))\n    \n    fig, axes = plt.subplots(1, len(indices), figsize=(20, 5))\n    if len(indices) == 1:\n        axes = [axes]  # Handle the case of a single image\n    \n    for i, idx in enumerate(indices):\n        axes[i].imshow(images[idx], cmap='gray')\n        axes[i].set_title(f'Slice: {filenames[idx]}')\n        axes[i].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_PATH, 'sample_slices.png'))\n    plt.show()\n    \n    return indices  # Return the indices for later use\n\n# Display and remember the sample indices for later comparison\nsample_indices = display_sample_slices(images, filenames)\n\n# 5. Load specific slices for training (8 slices distributed throughout the volume)\ndef load_specific_slices(file_list, num_slices=8):\n    \"\"\"Load specific slices from the file list, evenly distributed through the volume\"\"\"\n    total_files = len(file_list)\n    \n    if total_files < num_slices:\n        print(f\"Warning: Requested {num_slices} slices but only {total_files} available\")\n        indices = list(range(total_files))\n    else:\n        # Calculate evenly spaced indices\n        step = total_files // num_slices\n        indices = [i * step for i in range(num_slices)]\n        # Make sure we don't exceed the list bounds\n        indices = [min(i, total_files - 1) for i in indices]\n    \n    training_images = []\n    training_filenames = []\n    \n    for idx in indices:\n        file = file_list[idx]\n        img = np.array(Image.open(file).convert('L'))\n        training_images.append(img)\n        training_filenames.append(os.path.basename(file))\n    \n    print(f\"Loaded {len(training_images)} slices for training\")\n    print(f\"Training on slices at indices: {indices}\")\n    print(f\"Training on: {training_filenames}\")\n    \n    return training_images, training_filenames, indices\n\n# Load 8 evenly distributed slices for training\ntraining_images, training_filenames, training_indices = load_specific_slices(all_files, num_slices=8)\n\n# Display the training images\nplt.figure(figsize=(16, 8))\nrows = 2\ncols = (len(training_images) + 1) // 2  # Ceiling division for an odd number of images\n\nfor i, img in enumerate(training_images):\n    plt.subplot(rows, cols, i + 1)\n    plt.imshow(img, cmap='gray')\n    plt.title(f'Training: {training_filenames[i]}')\n    plt.axis('off')\n        \nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_PATH, 'training_images.png'))\nplt.show()\n\n# 6. Train on multiple slices\ndef train_n2v_model_multi(train_images):\n    \"\"\"Train a Noise2Void model on multiple images\"\"\"\n    datagen = N2V_DataGenerator()\n    \n    # Process each training image\n    all_patches = []\n    \n    for img in train_images:\n        # Prepare the image for N2V - add channel dimension\n        img_for_patches = img[np.newaxis, ..., np.newaxis]\n        \n        # Generate patches for this image\n        patches = datagen.generate_patches_from_list([img_for_patches], shape=(64, 64))\n        all_patches.append(patches)\n    \n    # Combine all patches\n    X = np.concatenate(all_patches, axis=0)\n    print(f\"Generated a total of {X.shape[0]} training patches of shape {X.shape[1:]}\")\n    \n    # Split into training and validation (80/20)\n    np.random.shuffle(X)\n    n_train = int(0.8 * X.shape[0])\n    X_train, X_val = X[:n_train], X[n_train:]\n    print(f\"Training on {n_train} patches, validating on {X.shape[0] - n_train} patches\")\n    \n    # Configure N2V\n    config = N2VConfig(\n        X_train,\n        unet_kern_size=3,\n        train_steps_per_epoch=max(int(X_train.shape[0]/128), 10), \n        train_epochs=15,  # 15 epochs as requested\n        train_loss='mse',\n        batch_norm=True,\n        train_batch_size=128,\n        n2v_perc_pix=0.198,\n        n2v_patch_shape=(64, 64),\n        n2v_manipulator='uniform_withCP',\n        n2v_neighborhood_radius=5,\n        # N2V2 improvements\n        blurpool=True,\n        skip_skipone=True,\n        unet_residual=False\n    )\n    \n    # Create and train the model\n    model_name = 'n2v_cryoET_8slices'\n    model = N2V(config, model_name, basedir=MODEL_PATH)\n    \n    print(\"Starting model training...\")\n    history = model.train(X_train, X_val)\n    \n    # Plot training history\n    plt.figure(figsize=(10, 4))\n    plot_history(history, ['loss', 'val_loss'])\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_PATH, 'training_history.png'))\n    plt.show()\n    \n    return model\n\n# Train the model using 8 images\nprint(f\"Training on {len(training_images)} slices with 15 epochs...\")\nmodel = train_n2v_model_multi(training_images)\n\n# 7. Denoise and compare the sample images\ndef process_and_compare(model, images, indices, filenames):\n    \"\"\"Process the selected slices and display before/after comparison\"\"\"\n    results = []\n    original_imgs = []\n    denoised_imgs = []\n    \n    for idx in indices:\n        # Get the original image\n        original = images[idx]\n        original_imgs.append(original)\n        \n        # Add dimensions for prediction: YXC\n        img_for_pred = original[..., np.newaxis]\n        \n        # Predict (denoise)\n        denoised = model.predict(img_for_pred, axes='YXC')\n        \n        # Remove the channel dimension for display\n        denoised = denoised[..., 0]\n        denoised_imgs.append(denoised)\n        \n        # Calculate noise metrics\n        noise_before = np.std(original)\n        noise_after = np.std(denoised)\n        results.append({\n            'filename': filenames[idx],\n            'noise_before': noise_before,\n            'noise_after': noise_after,\n            'reduction_percent': (1 - noise_after/noise_before) * 100\n        })\n    \n    # Display comparison\n    fig, axes = plt.subplots(2, len(indices), figsize=(20, 8))\n    \n    for i in range(len(original_imgs)):\n        # Original image\n        axes[0, i].imshow(original_imgs[i], cmap='gray')\n        axes[0, i].set_title(f'Original: {filenames[indices[i]]}')\n        axes[0, i].axis('off')\n        \n        # Denoised image\n        axes[1, i].imshow(denoised_imgs[i], cmap='gray')\n        axes[1, i].set_title(f'Denoised: {results[i][\"reduction_percent\"]:.1f}% reduction')\n        axes[1, i].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_PATH, 'denoising_comparison.png'))\n    plt.show()\n    \n    # Save individual images\n    for i, idx in enumerate(indices):\n        Image.fromarray(original_imgs[i].astype(np.uint8)).save(\n            os.path.join(OUTPUT_PATH, f'original_{filenames[indices[i]]}'))\n        Image.fromarray(denoised_imgs[i].astype(np.uint8)).save(\n            os.path.join(OUTPUT_PATH, f'denoised_{filenames[indices[i]]}'))\n    \n    return results, original_imgs, denoised_imgs\n\nprint(\"Denoising sample slices...\")\ndenoising_results, original_crops, denoised_crops = process_and_compare(model, images, sample_indices, filenames)\n\n# 8. Display detailed results\nfor result in denoising_results:\n    print(f\"Slice: {result['filename']}\")\n    print(f\"  Noise before: {result['noise_before']:.2f}\")\n    print(f\"  Noise after:  {result['noise_after']:.2f}\")\n    print(f\"  Reduction:    {result['reduction_percent']:.2f}%\")\n    print(\"\")\n\n# 9. Show detailed crop comparisons\ndef show_detailed_crops(original_imgs, denoised_imgs, indices, filenames):\n    \"\"\"Show detailed crops of the central region for better comparison\"\"\"\n    fig, axes = plt.subplots(2, len(indices), figsize=(20, 8))\n    \n    for i in range(len(original_imgs)):\n        # Create a central crop\n        h, w = original_imgs[i].shape\n        crop_y, crop_x = h//2 - 100, w//2 - 100\n        crop_h, crop_w = 200, 200\n        \n        crop_original = original_imgs[i][crop_y:crop_y+crop_h, crop_x:crop_x+crop_w]\n        crop_denoised = denoised_imgs[i][crop_y:crop_y+crop_h, crop_x:crop_x+crop_w]\n        \n        # Original crop\n        axes[0, i].imshow(crop_original, cmap='gray')\n        axes[0, i].set_title(f'Original: {filenames[indices[i]]} (crop)')\n        axes[0, i].axis('off')\n        \n        # Denoised crop\n        axes[1, i].imshow(crop_denoised, cmap='gray')\n        axes[1, i].set_title('Denoised (crop)')\n        axes[1, i].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_PATH, 'detailed_crop_comparison.png'))\n    plt.show()\n\nshow_detailed_crops(original_crops, denoised_crops, sample_indices, filenames)\n\n# 10. Process and denoise all slices (optional)\ndef process_all_slices():\n    \"\"\"Process and save denoised versions of all slices\"\"\"\n    # Ask user before proceeding\n    print(\"\\nDo you want to process and save denoised versions of all test slices?\")\n    print(\"This will save denoised versions of all slices loaded for testing.\")\n    process_all = input(\"Type 'yes' to proceed, or anything else to skip: \")\n    \n    if process_all.lower() != 'yes':\n        print(\"Skipping full dataset processing.\")\n        return\n    \n    print(f\"Processing all {len(images)} test slices...\")\n    output_dir = os.path.join(OUTPUT_PATH, 'all_denoised')\n    os.makedirs(output_dir, exist_ok=True)\n    \n    for i, (img, fname) in enumerate(tqdm(zip(images, filenames), total=len(images))):\n        # Add channel dimension for prediction\n        img_for_pred = img[..., np.newaxis]\n        \n        # Denoise the image\n        denoised = model.predict(img_for_pred, axes='YXC')[..., 0]\n        \n        # Save the denoised image\n        output_path = os.path.join(output_dir, f'denoised_{fname}')\n        Image.fromarray(denoised.astype(np.uint8)).save(output_path)\n    \n    print(f\"All slices processed and saved to {output_dir}\")\n\n# Uncomment to enable processing all slices\n# process_all_slices()\n\nprint(\"Noise2Void multi-slice training completed successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T22:40:28.398946Z","iopub.execute_input":"2025-03-26T22:40:28.399289Z","iopub.status.idle":"2025-03-26T23:06:25.448764Z","shell.execute_reply.started":"2025-03-26T22:40:28.399261Z","shell.execute_reply":"2025-03-26T23:06:25.447728Z"}},"outputs":[],"execution_count":null}]}