{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":91249,"databundleVersionId":11294684},{"sourceType":"datasetVersion","sourceId":12070702,"datasetId":6959173,"databundleVersionId":12596343}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Install and import libraries","metadata":{}},{"cell_type":"code","source":"pip install scikit-image scipy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:34.038476Z","iopub.execute_input":"2026-03-05T09:39:34.039067Z","iopub.status.idle":"2026-03-05T09:39:37.967881Z","shell.execute_reply.started":"2026-03-05T09:39:34.039038Z","shell.execute_reply":"2026-03-05T09:39:37.967147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2 as cv\nimport matplotlib.pyplot as plt\nfrom skimage.util import img_as_float\nfrom skimage.transform import resize\nfrom tensorflow.keras.losses import MeanSquaredError\nfrom tensorflow.keras.metrics import MeanAbsoluteError\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras import layers, models\nfrom skimage.io import imread, imsave\nimport tensorflow as tf\nimport numpy as np\nimport os\nfrom skimage import  io, transform, color, util\nfrom skimage.io import imread\nfrom scipy.ndimage import center_of_mass\nfrom scipy.ndimage import gaussian_filter,binary_opening, binary_closing\nimport os\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:37.969063Z","iopub.execute_input":"2026-03-05T09:39:37.969301Z","iopub.status.idle":"2026-03-05T09:39:51.205815Z","shell.execute_reply.started":"2026-03-05T09:39:37.969251Z","shell.execute_reply":"2026-03-05T09:39:51.205217Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Import Dataset","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train'\n\ntrain_labels = pd.read_csv('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.206809Z","iopub.execute_input":"2026-03-05T09:39:51.207304Z","iopub.status.idle":"2026-03-05T09:39:51.223307Z","shell.execute_reply.started":"2026-03-05T09:39:51.207274Z","shell.execute_reply":"2026-03-05T09:39:51.222624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_labels.iloc[1]\nsample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.224190Z","iopub.execute_input":"2026-03-05T09:39:51.224488Z","iopub.status.idle":"2026-03-05T09:39:51.244567Z","shell.execute_reply.started":"2026-03-05T09:39:51.224467Z","shell.execute_reply":"2026-03-05T09:39:51.243834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sorted(os.listdir(os.path.join(TRAIN_DIR,sample['tomo_id'])))[int(sample['Motor axis 0'])]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.246260Z","iopub.execute_input":"2026-03-05T09:39:51.246911Z","iopub.status.idle":"2026-03-05T09:39:51.265404Z","shell.execute_reply.started":"2026-03-05T09:39:51.246890Z","shell.execute_reply":"2026-03-05T09:39:51.264878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_image(sample, folder_path):\n    tomo_id = sample['tomo_id']\n    files = sorted(os.listdir(os.path.join(folder_path, tomo_id)))\n    file_name = files[int(sample['Motor axis 0'])]\n    img_path = os.path.join(folder_path, tomo_id, file_name)\n    \n    img = imread(img_path)\n\n    # If grayscale (2D), convert to 3 channels by stacking\n    if img.ndim == 2:\n        img = np.stack([img] * 3, axis=-1)  # (H, W) → (H, W, 3)\n\n    # If RGBA (4 channels), discard alpha\n    elif img.shape[-1] == 4:\n        img = img[..., :3]\n\n    return img\n\n\nsample = train_labels.iloc[2]\n\n\ndef generate_gaussian_heatmap(shape, center, radius):\n    \"\"\"\n    Generates a 2D Gaussian heatmap.\n\n    Parameters:\n        shape  : Tuple[int, int] - Shape of the output heatmap (height, width)\n        center : Tuple[int, int] - (x, y) coordinates of the center\n        radius : float           - Radius or standard deviation of the Gaussian\n\n    Returns:\n        heatmap: 2D NumPy array of shape `shape`\n    \"\"\"\n    x = np.arange(0, shape[1], 1)\n    y = np.arange(0, shape[0], 1)\n    xx, yy = np.meshgrid(x, y)\n\n    x0, y0 = center\n\n    # Gaussian formula\n    heatmap = np.exp(-((xx - x0) ** 2 + (yy - y0) ** 2) / (2 * radius ** 2))\n\n    # Normalize the heatmap to [0, 1]\n    heatmap = heatmap / np.max(heatmap)\n    return heatmap\n\nimg = read_image(sample,TRAIN_DIR)\n\n\n# If image is RGB, convert it to grayscale\n# if img.ndim == 3:\n#     img = color.rgb2gray(img)  # Returns float64 in [0, 1]\n\n# Optional: scale to 0–255 and convert to uint8 for consistent datatype\n# img = (img * 255).astype('uint8')\nheat_map = generate_gaussian_heatmap(img.shape,(int(sample['Motor axis 2']),int(sample['Motor axis 1'])),40)\nimg_copy = resize(\n    img,\n    (256,256),\n    anti_aliasing=True\n)\nheat_map = resize(\n    heat_map,\n    (256,256),\n    anti_aliasing=True\n)\nplt.figure(figsize=(10,10))\nplt.subplot(1,2,1)\nplt.imshow(img_copy)\nplt.subplot(1,2,2)\nplt.imshow(heat_map)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.265980Z","iopub.execute_input":"2026-03-05T09:39:51.266148Z","iopub.status.idle":"2026-03-05T09:39:51.790416Z","shell.execute_reply.started":"2026-03-05T09:39:51.266132Z","shell.execute_reply":"2026-03-05T09:39:51.789695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(img.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.791272Z","iopub.execute_input":"2026-03-05T09:39:51.791493Z","iopub.status.idle":"2026-03-05T09:39:51.795282Z","shell.execute_reply.started":"2026-03-05T09:39:51.791475Z","shell.execute_reply":"2026-03-05T09:39:51.794545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_mask(img,center,sigma=3.0,quantile_=0.70,length=40):\n\n    if img.ndim == 3:\n        img = np.mean(img, axis=-1).astype(np.uint8)  # Fast grayscale conversion\n\n    # print(img.shape)\n    mask = np.zeros_like(img)\n    mask1 = mask.copy()\n    mask2 = mask.copy()\n    \n    height = img.shape[0]\n    width = img.shape[1]\n    y = center[1]\n    x = center[0]\n    y1, y2 = max(0, y-length), min(height, y+length)\n    x1, x2 = max(0, x-length), min(width, x+length)\n    patch = img[y1:y2, x1:x2]\n    patch2 = patch.copy()\n\n    patch2 = gaussian_filter(patch2, sigma=sigma)\n    \n    q = np.quantile(patch,quantile_)\n    segmented_patch = patch <= q\n    mask[y1:y2,x1:x2] = segmented_patch\n    \n    q2 = np.quantile(patch2,quantile_)\n    segmented_patch2 = patch2 <= q2\n    mask1[y1:y2,x1:x2] = segmented_patch2\n    print(segmented_patch2.shape)\n\n    segmented_patch3 = segmented_patch2.copy()\n    segmented_patch3 = binary_opening(segmented_patch3, structure=np.ones(( 3, 3)))\n    segmented_patch3 = binary_closing(segmented_patch3, structure=np.ones(( 10, 10)))\n    mask2[y1:y2,x1:x2] = segmented_patch3\n    \n    return mask, mask1, mask2\n\ncenter = (int(sample['Motor axis 2']),int(sample['Motor axis 1']))\n\nmask1, mask2, mask3 = get_mask(img,center,length=60)\n\nplt.figure(figsize=(10,10))\nplt.subplot(2,2,1)\nplt.imshow(img)\nplt.title('sample')\nplt.subplot(2,2,2)\nplt.imshow(mask1)\nplt.title('mask')\nplt.subplot(2,2,3)\nplt.imshow(mask2)\nplt.title('mask with gaussin blur')\nplt.subplot(2,2,4)\nplt.imshow(mask3)\nplt.title('mask with morphing')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:51.796509Z","iopub.execute_input":"2026-03-05T09:39:51.796737Z","iopub.status.idle":"2026-03-05T09:39:52.813026Z","shell.execute_reply.started":"2026-03-05T09:39:51.796718Z","shell.execute_reply":"2026-03-05T09:39:52.812319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels_2 = train_labels[train_labels['Number of motors'] > 0]\ntrain_labels_2.describe()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:52.813838Z","iopub.execute_input":"2026-03-05T09:39:52.814089Z","iopub.status.idle":"2026-03-05T09:39:52.852384Z","shell.execute_reply.started":"2026-03-05T09:39:52.814069Z","shell.execute_reply":"2026-03-05T09:39:52.851633Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.hist(train_labels_2['Number of motors'],bins=10,rwidth=0.8)\nplt.ylabel(\"No of samples\")\nplt.xlabel(\"No of motors\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:52.853230Z","iopub.execute_input":"2026-03-05T09:39:52.853531Z","iopub.status.idle":"2026-03-05T09:39:53.121494Z","shell.execute_reply.started":"2026-03-05T09:39:52.853507Z","shell.execute_reply":"2026-03-05T09:39:53.120743Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"DATASET_DIR = '/kaggle/working/dataset'\nos.makedirs(DATASET_DIR,exist_ok=True)\nos.makedirs(os.path.join(DATASET_DIR,\"train\"),exist_ok=True)\nos.makedirs(os.path.join(DATASET_DIR,\"test\"),exist_ok=True)\nos.makedirs(os.path.join(DATASET_DIR,\"val\"),exist_ok=True)\nTRAIN_DATASET_DIR = os.path.join(DATASET_DIR,\"train\")\nTEST_DATASET_DIR = os.path.join(DATASET_DIR,\"test\")\nVAL_DATASET_DIR = os.path.join(DATASET_DIR,\"val\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.124150Z","iopub.execute_input":"2026-03-05T09:39:53.124445Z","iopub.status.idle":"2026-03-05T09:39:53.129634Z","shell.execute_reply.started":"2026-03-05T09:39:53.124426Z","shell.execute_reply":"2026-03-05T09:39:53.128978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXTRA_DATASET_DIR = '/kaggle/input/cryoet-flagellar-motors-dataset/jpgs'\nextra_dataset_labels = pd.read_csv('/kaggle/input/cryoet-flagellar-motors-dataset/labels.csv')\nextra_dataset_labels_new = pd.read_csv('/kaggle/input/cryoet-flagellar-motors-dataset/labels_new.csv')\ndisplay(extra_dataset_labels_new.info())\nprint(f'No of samples in extra dataset: {len(os.listdir(EXTRA_DATASET_DIR))}')\nfiles_in_dir = os.listdir(EXTRA_DATASET_DIR)\nmissing_samples = extra_dataset_labels_new[~extra_dataset_labels_new['tomo_id'].isin(files_in_dir)]\n\nprint(f\"No of samples missing in dir: {len(missing_samples)}\")\ndisplay(extra_dataset_labels_new)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.130388Z","iopub.execute_input":"2026-03-05T09:39:53.130582Z","iopub.status.idle":"2026-03-05T09:39:53.200298Z","shell.execute_reply.started":"2026-03-05T09:39:53.130561Z","shell.execute_reply":"2026-03-05T09:39:53.199516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"clean_extra_data = extra_dataset_labels_new[extra_dataset_labels_new['tomo_id'].isin(files_in_dir)]\nclean_extra_data = clean_extra_data[(clean_extra_data['z']>0) & (clean_extra_data['y']>0) & (clean_extra_data['x']>0)]\ndisplay(clean_extra_data.info())\ndisplay(clean_extra_data.describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.201494Z","iopub.execute_input":"2026-03-05T09:39:53.201757Z","iopub.status.idle":"2026-03-05T09:39:53.233284Z","shell.execute_reply.started":"2026-03-05T09:39:53.201740Z","shell.execute_reply":"2026-03-05T09:39:53.232697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def generate_mask(sample,folder_path , dataset_dir=TRAIN_DIR):\n#     tomo_id = sample['tomo_id']\n#     mask_folder = os.path.join(folder_path,'mask')\n#     mask_file = os.path.join(mask_folder,tomo_id,f\"slice_{sample['Motor axis 0']}.png\")\n#     if os.path.exist(mask_file):\n        \n\n# def flagellar_dataset(samples,dataset_dir):\n#     for idx, sample in samples.iterrows():\n#         img = read_image(sample,TRAIN_DIR)\n        \n\ndef get_custom_mask(img,center,sigma=3.0,quantile_=0.70,length=40):\n\n    if img.ndim == 3:\n        img = np.mean(img, axis=-1).astype(np.uint8)  # Fast grayscale conversion\n\n    # print(img.shape)\n    mask = np.zeros_like(img)\n    mask1 = mask.copy()\n    mask2 = mask.copy()\n    \n    height = img.shape[0]\n    width = img.shape[1]\n    y = center[1]\n    x = center[0]\n    y1, y2 = max(0, y-length), min(height, y+length)\n    x1, x2 = max(0, x-length), min(width, x+length)\n    patch = img[y1:y2, x1:x2]\n    patch2 = patch.copy()\n\n    patch2 = gaussian_filter(patch2, sigma=sigma)\n    \n    q = np.quantile(patch,quantile_)\n    segmented_patch = patch <= q\n    mask[y1:y2,x1:x2] = segmented_patch\n    \n    q2 = np.quantile(patch2,quantile_)\n    segmented_patch2 = patch2 <= q2\n    mask1[y1:y2,x1:x2] = segmented_patch2\n    # print(segmented_patch2.shape)\n\n    segmented_patch3 = segmented_patch2.copy()\n    segmented_patch3 = binary_opening(segmented_patch3, structure=np.ones(( 3, 3)))\n    segmented_patch3 = binary_closing(segmented_patch3, structure=np.ones(( 10, 10)))\n    mask2[y1:y2,x1:x2] = segmented_patch3\n    \n    return mask2\n\n# --------------------- Gaussian Heatmap --------------------- #\ndef generate_gaussian_heatmap(shape, center, radius):\n    x = np.arange(0, shape[1], 1)\n    y = np.arange(0, shape[0], 1)\n    xx, yy = np.meshgrid(x, y)\n    x0, y0 = center\n    heatmap = np.exp(-((xx - x0) ** 2 + (yy - y0) ** 2) / (2 * radius ** 2))\n    heatmap = heatmap / np.max(heatmap)\n    return heatmap\n\n# --------------------- Read Image --------------------- #\ndef read_image(sample, folder_path):\n    files = sorted(os.listdir(os.path.join(folder_path, sample['tomo_id'])))\n    file_name = files[int(sample['Motor axis 0'])]\n    return imread(os.path.join(folder_path, sample['tomo_id'], file_name))\n\n# --------------------- Generate Mask & Save --------------------- #\ndef generate_mask(sample, dataset_dir, output_base_dir, radius=60):\n    tomo_id = sample['tomo_id']\n    slice_idx = int(sample['Motor axis 0'])\n    center_x = float(sample['Motor axis 2'])\n    center_y = float(sample['Motor axis 1'])\n\n    # Read image\n    image = read_image(sample, dataset_dir)\n    image_shape = image.shape[:2]\n\n    # # 1. Save image to output/img/{tomo_id}/slice_{idx}.png\n    img_dir = os.path.join(output_base_dir, 'img')\n    os.makedirs(img_dir, exist_ok=True)\n    img_path = os.path.join(output_base_dir, 'img', f\"{ tomo_id}_slice_{slice_idx}.png\")\n    if not os.path.exists(img_path):  # avoid overwriting\n        imsave(img_path, image)\n\n    # 2. Create Gaussian heatmap\n    # heatmap = generate_gaussian_heatmap(image_shape, center=(center_x, center_y), radius=radius)\n\n    # get quantile mask\n    heatmap = get_custom_mask(image,(int(center_x), int(center_y)))\n    \n    # # 3. Save mask to output/mask/{tomo_id}/slice_{idx}.png\n    mask_dir = os.path.join(output_base_dir, 'mask')\n    os.makedirs(mask_dir, exist_ok=True)\n    mask_path = os.path.join(output_base_dir, 'mask',  f\"{tomo_id}_slice_{slice_idx}.png\")\n\n    if os.path.exists(mask_path):\n        # Load existing and add heatmap\n        existing_mask = imread(mask_path).astype(np.float32) / 255.0\n        combined_mask = np.clip(existing_mask + heatmap, 0, 1)\n    else:\n        combined_mask = heatmap\n\n    imsave(mask_path, (combined_mask * 255).astype(np.uint8))\n\n# --------------------- Process All Samples --------------------- #\ndef flagellar_dataset(samples, output_base_dir,dataset_dir=TRAIN_DIR):\n    for idx, sample in samples.iterrows():\n        generate_mask(sample, dataset_dir=dataset_dir, output_base_dir=output_base_dir)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.233995Z","iopub.execute_input":"2026-03-05T09:39:53.234247Z","iopub.status.idle":"2026-03-05T09:39:53.245839Z","shell.execute_reply.started":"2026-03-05T09:39:53.234222Z","shell.execute_reply":"2026-03-05T09:39:53.245230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"No of samples: {len(train_labels_2['tomo_id'])}\")\ntrain_split = int(len(train_labels_2['tomo_id'])*0.70)\ntest_split = int(len(train_labels_2['tomo_id'])*0.15)\ntrain_samples = train_labels_2.iloc[:train_split]\ntest_samples = train_labels_2.iloc[train_split:train_split+test_split]\nval_samples = train_labels_2.iloc[train_split+test_split:]\nprint(f\"No of train samples: {len(train_samples['tomo_id'])}\")\nprint(f\"No of tset samples: {len(test_samples['tomo_id'])}\")\nprint(f\"No of val samples: {len(val_samples['tomo_id'])}\")\n# flagellar_dataset(train_labels_2.iloc[:4],)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.246473Z","iopub.execute_input":"2026-03-05T09:39:53.246726Z","iopub.status.idle":"2026-03-05T09:39:53.259798Z","shell.execute_reply.started":"2026-03-05T09:39:53.246710Z","shell.execute_reply":"2026-03-05T09:39:53.259191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(train_samples,output_base_dir=TRAIN_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:39:53.260492Z","iopub.execute_input":"2026-03-05T09:39:53.260720Z","iopub.status.idle":"2026-03-05T09:40:27.899964Z","shell.execute_reply.started":"2026-03-05T09:39:53.260694Z","shell.execute_reply":"2026-03-05T09:40:27.899189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(test_samples,output_base_dir=TEST_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:27.900703Z","iopub.execute_input":"2026-03-05T09:40:27.900907Z","iopub.status.idle":"2026-03-05T09:40:35.413700Z","shell.execute_reply.started":"2026-03-05T09:40:27.900890Z","shell.execute_reply":"2026-03-05T09:40:35.413078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(val_samples,output_base_dir=VAL_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:35.414389Z","iopub.execute_input":"2026-03-05T09:40:35.414662Z","iopub.status.idle":"2026-03-05T09:40:42.681277Z","shell.execute_reply.started":"2026-03-05T09:40:35.414640Z","shell.execute_reply":"2026-03-05T09:40:42.680651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = '/kaggle/working/dataset/test/img'\ntrain_dir = '/kaggle/working/dataset/train/img'\nval_dir = '/kaggle/working/dataset/val/img'\nprint(f\"No of train samples:{len(os.listdir(train_dir))}\")\nprint(f\"No of test samples:{len(os.listdir(test_dir))}\")\nprint(f\"No of val samples:{len(os.listdir(val_dir))}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.682117Z","iopub.execute_input":"2026-03-05T09:40:42.682594Z","iopub.status.idle":"2026-03-05T09:40:42.687487Z","shell.execute_reply.started":"2026-03-05T09:40:42.682563Z","shell.execute_reply":"2026-03-05T09:40:42.686876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels_2[(train_labels_2['Number of motors']==4) & (train_labels_2['tomo_id']=='tomo_1b82d1')]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.688206Z","iopub.execute_input":"2026-03-05T09:40:42.688890Z","iopub.status.idle":"2026-03-05T09:40:42.704074Z","shell.execute_reply.started":"2026-03-05T09:40:42.688871Z","shell.execute_reply":"2026-03-05T09:40:42.703542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels_2[(train_labels_2['Number of motors']==10)& (train_labels_2['tomo_id']=='tomo_226cd8')]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.704721Z","iopub.execute_input":"2026-03-05T09:40:42.704959Z","iopub.status.idle":"2026-03-05T09:40:42.720892Z","shell.execute_reply.started":"2026-03-05T09:40:42.704939Z","shell.execute_reply":"2026-03-05T09:40:42.720235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"clean_extra_data = clean_extra_data.rename(columns={'z':'Motor axis 0','y':'Motor axis 1','x':'Motor axis 2'})\nclean_extra_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.721646Z","iopub.execute_input":"2026-03-05T09:40:42.722650Z","iopub.status.idle":"2026-03-05T09:40:42.737705Z","shell.execute_reply.started":"2026-03-05T09:40:42.722631Z","shell.execute_reply":"2026-03-05T09:40:42.737131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"extra_data_train_split = int(len(clean_extra_data['tomo_id'])*0.80)\nextra_data_test_split = int(len(clean_extra_data['tomo_id'])*0.10)\nextra_train_data = clean_extra_data.iloc[:extra_data_train_split]\nextra_test_data = clean_extra_data.iloc[extra_data_train_split:extra_data_train_split+extra_data_test_split]\nextra_val_data = clean_extra_data.iloc[extra_data_train_split+extra_data_test_split:]\nprint(f\"no of extra train samples: {len(extra_train_data['tomo_id'])}\")\nprint(f\"no of extra test samples: {len(extra_test_data['tomo_id'])}\")\nprint(f\"no of extra val samples: {len(extra_val_data['tomo_id'])}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.738396Z","iopub.execute_input":"2026-03-05T09:40:42.738679Z","iopub.status.idle":"2026-03-05T09:40:42.749562Z","shell.execute_reply.started":"2026-03-05T09:40:42.738663Z","shell.execute_reply":"2026-03-05T09:40:42.748972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(extra_train_data,output_base_dir=TRAIN_DATASET_DIR,dataset_dir=EXTRA_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:40:42.752237Z","iopub.execute_input":"2026-03-05T09:40:42.752527Z","iopub.status.idle":"2026-03-05T09:41:30.194273Z","shell.execute_reply.started":"2026-03-05T09:40:42.752510Z","shell.execute_reply":"2026-03-05T09:41:30.193715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(extra_test_data,output_base_dir=TEST_DATASET_DIR,dataset_dir=EXTRA_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:30.194982Z","iopub.execute_input":"2026-03-05T09:41:30.195190Z","iopub.status.idle":"2026-03-05T09:41:35.967175Z","shell.execute_reply.started":"2026-03-05T09:41:30.195174Z","shell.execute_reply":"2026-03-05T09:41:35.966615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"flagellar_dataset(extra_val_data,output_base_dir=VAL_DATASET_DIR,dataset_dir=EXTRA_DATASET_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:35.967924Z","iopub.execute_input":"2026-03-05T09:41:35.968128Z","iopub.status.idle":"2026-03-05T09:41:41.374987Z","shell.execute_reply.started":"2026-03-05T09:41:35.968111Z","shell.execute_reply":"2026-03-05T09:41:41.374186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = '/kaggle/working/dataset/test/img'\ntrain_dir = '/kaggle/working/dataset/train/img'\nval_dir = '/kaggle/working/dataset/val/img'\nprint(f\"No of training samples: {len(os.listdir(train_dir))}\")\nprint(f\"No of test samples: {len(os.listdir(test_dir))}\")\nprint(f\"No of val samples: {len(os.listdir(val_dir))}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:41.375846Z","iopub.execute_input":"2026-03-05T09:41:41.376052Z","iopub.status.idle":"2026-03-05T09:41:41.382177Z","shell.execute_reply.started":"2026-03-05T09:41:41.376035Z","shell.execute_reply":"2026-03-05T09:41:41.381488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = (256, 256)\nMASK_SIZE = (256, 256)\ndef load_image_and_mask(img_path, mask_path):\n    img = tf.io.read_file(img_path)\n    img = tf.image.decode_png(img, channels=3)\n    img = tf.image.resize(img, IMG_SIZE)\n    img = tf.cast(img, tf.float32) / 255.0\n\n    mask = tf.io.read_file(mask_path)\n    mask = tf.image.decode_png(mask, channels=1)\n    mask = tf.image.resize(mask, MASK_SIZE)\n    mask = tf.cast(mask, tf.float32) / 255.0  # normalize heatmap\n\n    return img, mask\n\ndef create_heatmap_dataset(image_dir, mask_dir, batch_size=16, shuffle=True):\n    image_files = sorted([f for f in os.listdir(image_dir) if f.endswith('.png')])\n    mask_files = sorted([f for f in os.listdir(mask_dir) if f.endswith('.png')])\n\n    image_paths = [os.path.join(image_dir, f) for f in image_files]\n    mask_paths = [os.path.join(mask_dir, f) for f in mask_files]\n\n    def gen():\n        for img_path, mask_path in zip(image_paths, mask_paths):\n            yield load_image_and_mask(img_path, mask_path)\n\n    dataset = tf.data.Dataset.from_generator(\n        gen,\n        output_signature=(\n            tf.TensorSpec(shape=(256, 256, 3), dtype=tf.float32),\n            tf.TensorSpec(shape=(256, 256, 1), dtype=tf.float32),\n        )\n    )\n\n    if shuffle:\n        dataset = dataset.shuffle(buffer_size=1024)\n    return dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:41.383061Z","iopub.execute_input":"2026-03-05T09:41:41.383295Z","iopub.status.idle":"2026-03-05T09:41:41.395626Z","shell.execute_reply.started":"2026-03-05T09:41:41.383254Z","shell.execute_reply":"2026-03-05T09:41:41.394906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import tensorflow as tf\n# import os\n\n# IMG_SIZE = (256, 256)\n\n# def load_image_and_mask(img_path, mask_path):\n#     img = tf.io.read_file(img_path)\n#     img = tf.image.decode_png(img, channels=3)\n#     img = tf.image.resize(img, IMG_SIZE)\n#     img = tf.cast(img, tf.float32) / 255.0\n\n#     mask = tf.io.read_file(mask_path)\n#     mask = tf.image.decode_png(mask, channels=1)\n#     mask = tf.image.resize(mask, IMG_SIZE)\n#     mask = tf.cast(mask, tf.float32) / 255.0\n\n#     return img, mask\n\n# def augment(img, mask):\n#     if tf.random.uniform(()) > 0.5:\n#         img = tf.image.flip_left_right(img)\n#         mask = tf.image.flip_left_right(mask)\n#     if tf.random.uniform(()) > 0.5:\n#         img = tf.image.flip_up_down(img)\n#         mask = tf.image.flip_up_down(mask)\n#     k = tf.random.uniform(shape=[], minval=0, maxval=4, dtype=tf.int32)\n#     img = tf.image.rot90(img, k=k)\n#     mask = tf.image.rot90(mask, k=k)\n#     img = tf.image.random_brightness(img, max_delta=0.2)\n#     img = tf.image.random_contrast(img, lower=0.8, upper=1.2)\n#     return img, mask\n\n# def create_augmented_dataset(image_dir, mask_dir, batch_size=16, augment_factor=5, shuffle=True):\n#     image_files = sorted([f for f in os.listdir(image_dir) if f.endswith('.png')])\n#     mask_files = sorted([f for f in os.listdir(mask_dir) if f.endswith('.png')])\n\n#     image_paths = [os.path.join(image_dir, f) for f in image_files]\n#     mask_paths = [os.path.join(mask_dir, f) for f in mask_files]\n\n#     # Repeat each file path augment_factor times\n#     image_paths = image_paths * augment_factor\n#     mask_paths = mask_paths * augment_factor\n\n#     dataset = tf.data.Dataset.from_tensor_slices((image_paths, mask_paths))\n    \n#     if shuffle:\n#         dataset = dataset.shuffle(buffer_size=1024)\n\n#     # Load and augment\n#     dataset = dataset.map(load_image_and_mask, num_parallel_calls=tf.data.AUTOTUNE)\n#     dataset = dataset.map(augment, num_parallel_calls=tf.data.AUTOTUNE)\n\n#     return dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n\n\n\nimport tensorflow as tf\nimport os\n\nIMG_SIZE = (256, 256)\n\ndef load_image_and_mask(img_path, mask_path):\n    img = tf.io.read_file(img_path)\n    img = tf.image.decode_png(img, channels=3)\n    img = tf.image.resize(img, IMG_SIZE)\n    img = tf.cast(img, tf.float32) / 255.0\n\n    mask = tf.io.read_file(mask_path)\n    mask = tf.image.decode_png(mask, channels=1)\n    mask = tf.image.resize(mask, IMG_SIZE)\n    mask = tf.cast(mask, tf.float32) / 255.0\n\n    return img, mask\n\ndef augment(img, mask):\n    # Random horizontal flip\n    if tf.random.uniform(()) > 0.5:\n        img = tf.image.flip_left_right(img)\n        mask = tf.image.flip_left_right(mask)\n\n    # Random vertical flip\n    if tf.random.uniform(()) > 0.5:\n        img = tf.image.flip_up_down(img)\n        mask = tf.image.flip_up_down(mask)\n\n    # Random rotation\n    k = tf.random.uniform(shape=[], minval=0, maxval=4, dtype=tf.int32)\n    img = tf.image.rot90(img, k)\n    mask = tf.image.rot90(mask, k)\n\n    # Color-only augmentations\n    img = tf.image.random_brightness(img, max_delta=0.2)\n    img = tf.image.random_contrast(img, lower=0.8, upper=1.2)\n    img = tf.image.random_saturation(img, lower=0.8, upper=1.2)\n    img = tf.clip_by_value(img, 0.0, 1.0)\n\n    return img, mask\n\ndef create_augmented_dataset(image_dir, mask_dir, batch_size=16, augment_multiplier=5, shuffle=True):\n    \"\"\"\n    augment_multiplier: how many times to repeat each image with different augmentations.\n    \"\"\"\n    image_files = sorted([f for f in os.listdir(image_dir) if f.endswith('.png')])\n    mask_files = sorted([f for f in os.listdir(mask_dir) if f.endswith('.png')])\n\n    image_paths = [os.path.join(image_dir, f) for f in image_files]\n    mask_paths = [os.path.join(mask_dir, f) for f in mask_files]\n\n    base_dataset = tf.data.Dataset.from_tensor_slices((image_paths, mask_paths))\n\n    if shuffle:\n        base_dataset = base_dataset.shuffle(buffer_size=1024)\n\n    # Expand dataset by repeating and augmenting\n    augmented_dataset = base_dataset.map(\n        lambda x, y: tf.py_function(load_image_and_mask, [x, y], [tf.float32, tf.float32]),\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n\n    # Set shapes for TensorFlow graph\n    augmented_dataset = augmented_dataset.map(\n        lambda x, y: (tf.ensure_shape(x, (256, 256, 3)), tf.ensure_shape(y, (256, 256, 1)))\n    )\n\n    # Apply augmentation multiple times\n    datasets = []\n    for _ in range(augment_multiplier):\n        ds = augmented_dataset.map(augment, num_parallel_calls=tf.data.AUTOTUNE)\n        datasets.append(ds)\n\n    final_dataset = tf.data.Dataset.concatenate(datasets[0], datasets[1]) if augment_multiplier > 1 else datasets[0]\n    for i in range(2, augment_multiplier):\n        final_dataset = tf.data.Dataset.concatenate(final_dataset, datasets[i])\n\n    final_dataset = final_dataset.shuffle(1000).batch(batch_size).prefetch(tf.data.AUTOTUNE)\n    return final_dataset\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:41.396505Z","iopub.execute_input":"2026-03-05T09:41:41.396783Z","iopub.status.idle":"2026-03-05T09:41:41.411070Z","shell.execute_reply.started":"2026-03-05T09:41:41.396766Z","shell.execute_reply":"2026-03-05T09:41:41.410464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.layers import Input, Conv2D, UpSampling2D, Concatenate, Activation, BatchNormalization\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.applications import EfficientNetB0, ResNet50\nimport numpy as np\ninputs = Input((256, 256, 3))\n# base_model = ResNet50(weights='imagenet', include_top=False, input_tensor=inputs)\n# base_model.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:41.412505Z","iopub.execute_input":"2026-03-05T09:41:41.412729Z","iopub.status.idle":"2026-03-05T09:41:41.428781Z","shell.execute_reply.started":"2026-03-05T09:41:41.412710Z","shell.execute_reply":"2026-03-05T09:41:41.428189Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.layers import Input, Conv2D, UpSampling2D, Concatenate, Activation, BatchNormalization, Dropout\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.applications import EfficientNetB0, ResNet50\nfrom tensorflow.keras import layers\ndef unet_model_final_with_dropout(input_shape, backbone='EfficientNetB0', dropout_rate=0.3):\n\n    inputs = Input(input_shape)\n    IMG_SIZE = input_shape[0]\n    resize_layer = layers.Resizing(IMG_SIZE, IMG_SIZE)(inputs)\n    # --- Encoder Path ---\n    if backbone == 'EfficientNetB0':\n        base_model = EfficientNetB0(weights='imagenet', include_top=False, input_tensor=resize_layer)\n        skip_layers = [\n            'block2a_expand_activation',\n            'block3a_expand_activation',\n            'block4a_expand_activation',\n            'block6a_expand_activation',\n            'top_activation'\n        ]\n        skip_connections = [base_model.get_layer(name).output for name in skip_layers]\n        encoder_output = skip_connections[-1]\n        skip_connections = skip_connections[:-1]\n    \n    elif backbone == 'ResNet50':\n        base_model = ResNet50(weights='imagenet', include_top=False, input_tensor=inputs)\n        skip_layers = [\n            'conv2_block3_out',\n            'conv3_block4_out',\n            'conv4_block6_out',\n            'conv5_block3_out'\n        ]\n        skip_connections = [base_model.get_layer(name).output for name in skip_layers]\n        encoder_output = skip_connections[-1]\n        skip_connections = skip_connections[:-1]\n    else:\n        raise ValueError(\"Unsupported backbone. Choose 'EfficientNetB0' or 'ResNet50'.\")\n\n    # --- Decoder Path with Dropout ---\n    \n    # Bottleneck upsampling\n    d1 = UpSampling2D(size=(2, 2))(encoder_output)\n    d1 = Conv2D(512, (5, 5), padding='same')(d1)\n    d1 = Concatenate()([d1, skip_connections[-1]])\n    d1 = BatchNormalization()(d1)\n    d1 = Activation('relu')(d1)\n    d1 = Conv2D(512, (5, 5), padding='same')(d1)\n    d1 = BatchNormalization()(d1)\n    d1 = Activation('relu')(d1)\n    d1 = Dropout(dropout_rate)(d1) \n\n    # Decoder 2\n    d2 = UpSampling2D(size=(2, 2))(d1)\n    d2 = Conv2D(256, (5, 5), padding='same')(d2)\n    d2 = Concatenate()([d2, skip_connections[-2]])\n    d2 = BatchNormalization()(d2)\n    d2 = Activation('relu')(d2)\n    d2 = Conv2D(256, (5, 5), padding='same')(d2)\n    d2 = BatchNormalization()(d2)\n    d2 = Activation('relu')(d2)\n    d2 = Dropout(dropout_rate)(d2)\n\n    # Decoder 3\n    d3 = UpSampling2D(size=(2, 2))(d2)\n    d3 = Conv2D(128, (5, 5), padding='same')(d3)\n    d3 = Concatenate()([d3, skip_connections[-3]])\n    d3 = BatchNormalization()(d3)\n    d3 = Activation('relu')(d3)\n    d3 = Conv2D(128, (5, 5), padding='same')(d3)\n    d3 = BatchNormalization()(d3)\n    d3 = Activation('relu')(d3)\n    d3 = Dropout(dropout_rate)(d3) \n\n    # Decoder 4\n    if backbone == 'EfficientNetB0':\n        d4 = UpSampling2D(size=(2, 2))(d3)\n        d4 = Conv2D(64, (5, 5), padding='same')(d4)\n        d4 = Concatenate()([d4, skip_connections[-4]])\n        d4 = BatchNormalization()(d4)\n        d4 = Activation('relu')(d4)\n        d4 = Conv2D(64, (5, 5), padding='same')(d4)\n        d4 = BatchNormalization()(d4)\n        d4 = Activation('relu')(d4)\n        d4 = Dropout(dropout_rate)(d4) \n        final_conv = d4\n    else: \n        d4 = UpSampling2D(size=(2, 2))(d3)\n        d4 = Conv2D(64, (5, 5), padding='same')(d4)\n        d4 = Concatenate()([d4, skip_connections[-4]])\n        d4 = BatchNormalization()(d4)\n        d4 = Activation('relu')(d4)\n        d4 = Conv2D(64, (5, 5), padding='same')(d4)\n        d4 = BatchNormalization()(d4)\n        d4 = Activation('relu')(d4)\n        d4 = Dropout(dropout_rate)(d4) \n        final_conv = d4\n    \n    # Final upsampling to 256x256\n    final_upsample = UpSampling2D(size=(2, 2))(final_conv)\n    final_upsample = Conv2D(32, (5, 5), padding='same')(final_upsample)\n    final_upsample = BatchNormalization()(final_upsample)\n    final_upsample = Activation('relu')(final_upsample)\n    final_upsample = Dropout(dropout_rate)(final_upsample)\n\n    # Output layer\n    outputs = Conv2D(1, (1, 1), padding='same', activation='sigmoid')(final_upsample)\n\n    model = Model(inputs, outputs, name='unet_with_' + backbone + '_backbone')\n    \n    return model\n\n# Example usage with corrected code\ntry:\n    model = unet_model_final_with_dropout(input_shape=(256, 256, 3), backbone='EfficientNetB0', dropout_rate=0.4)\n    model.summary()\nexcept Exception as e:\n    print(f\"An error occurred: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:41.429518Z","iopub.execute_input":"2026-03-05T09:41:41.429765Z","iopub.status.idle":"2026-03-05T09:41:45.208694Z","shell.execute_reply.started":"2026-03-05T09:41:41.429749Z","shell.execute_reply":"2026-03-05T09:41:45.207976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = create_augmented_dataset('/kaggle/working/dataset/train/img', '/kaggle/working/dataset/train/mask',augment_multiplier=5)\ntest_ds = create_augmented_dataset('/kaggle/working/dataset/test/img', '/kaggle/working/dataset/test/mask')\nval_ds = create_augmented_dataset('/kaggle/working/dataset/val/img', '/kaggle/working/dataset/val/mask',augment_multiplier=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:45.209455Z","iopub.execute_input":"2026-03-05T09:41:45.209667Z","iopub.status.idle":"2026-03-05T09:41:48.210809Z","shell.execute_reply.started":"2026-03-05T09:41:45.209649Z","shell.execute_reply":"2026-03-05T09:41:48.210236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow.keras.backend as K\nfrom keras.saving import register_keras_serializable\n@register_keras_serializable()\ndef dice_coefficient(y_true, y_pred, smooth=1.0):\n    \"\"\"\n    Calculates the Dice Coefficient (Sørensen–Dice coefficient) as a metric.\n    \"\"\"\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = K.sum(y_true_f * y_pred_f)\n    return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n@register_keras_serializable()\ndef dice_loss(y_true, y_pred):\n    \"\"\"\n    Calculates the Dice Loss.\n    \"\"\"\n    return 1 - dice_coefficient(y_true, y_pred)\n\n@register_keras_serializable()\ndef iou_score(y_true, y_pred, smooth=1.0):\n    \"\"\"\n    Calculates the Intersection over Union (IoU) score as a metric.\n    \"\"\"\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = K.sum(y_true_f * y_pred_f)\n    union = K.sum(y_true_f) + K.sum(y_pred_f) - intersection\n    return (intersection + smooth) / (union + smooth)\n\n\nimport tensorflow.keras.backend as K\n@register_keras_serializable()\ndef tversky_loss(y_true, y_pred, alpha=0.3, beta=0.7, smooth=1e-6):\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n\n    TP = K.sum(y_true_f * y_pred_f)\n    FP = K.sum((1 - y_true_f) * y_pred_f)\n    FN = K.sum(y_true_f * (1 - y_pred_f))\n\n    return 1 - (TP + smooth) / (TP + alpha * FP + beta * FN + smooth)\n@register_keras_serializable()\ndef dice_coef(y_true, y_pred, smooth=1e-6):\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = K.sum(y_true_f * y_pred_f)\n    return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n\n\n# @register_keras_serializable()\n# def dice_coef_edema(y_true, y_pred, epsilon=1e-6):\n#     # Assuming edema is the 3rd channel (index 2)\n#     y_true_edema = y_true[:, :, :, 0]\n#     y_pred_edema = y_pred[:, :, :, 0]\n\n#     intersection = K.sum(K.abs(y_true_edema * y_pred_edema))\n#     denominator = K.sum(K.square(y_true_edema)) + K.sum(K.square(y_pred_edema)) + epsilon\n\n#     return (2. * intersection) / denominator\n\n\n@register_keras_serializable()\ndef combined_loss_penalizing_dice_and_sensitivity(y_true, y_pred, alpha=0.3, beta=0.7, lambda_dice=1.0):\n    # Global Tversky loss to focus on sensitivity\n    tversky = tversky_loss(y_true, y_pred, alpha=alpha, beta=beta)\n\n    # Dice loss for whole mask\n    # dice_1 = dice_coef_edema(y_true, y_pred)\n    dice_2 = dice_coef(y_true, y_pred)\n    # dice_loss_1 = 1 - dice_1\n    dice_loss_2 = 1 - dice_2\n\n    # Combine\n    return tversky + lambda_dice * (dice_loss_2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:48.211673Z","iopub.execute_input":"2026-03-05T09:41:48.211914Z","iopub.status.idle":"2026-03-05T09:41:48.221954Z","shell.execute_reply.started":"2026-03-05T09:41:48.211888Z","shell.execute_reply":"2026-03-05T09:41:48.221295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create the model\n# input_size = (256, 256, 3)\n# model = unet_model_revised(input_shape=input_size, backbone='EfficientNetB0')\n\n# Compile the model\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    loss=combined_loss_penalizing_dice_and_sensitivity, \n    metrics=[iou_score, dice_coefficient, 'accuracy']\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:48.222702Z","iopub.execute_input":"2026-03-05T09:41:48.223350Z","iopub.status.idle":"2026-03-05T09:41:48.274816Z","shell.execute_reply.started":"2026-03-05T09:41:48.223331Z","shell.execute_reply":"2026-03-05T09:41:48.273914Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"callbacks = [\n    tf.keras.callbacks.ModelCheckpoint(\n        filepath='best_model.h5',\n        monitor='val_loss',\n        save_best_only=True,\n        verbose=1\n    ),\n    tf.keras.callbacks.EarlyStopping(\n        monitor='val_loss',\n        patience=10,\n        restore_best_weights=True\n    )\n]\n\n# Train the model\nhistory = model.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=50,\n    callbacks=callbacks\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T09:41:48.275597Z","iopub.execute_input":"2026-03-05T09:41:48.275834Z","iopub.status.idle":"2026-03-05T10:06:40.772379Z","shell.execute_reply.started":"2026-03-05T09:41:48.275806Z","shell.execute_reply":"2026-03-05T10:06:40.771310Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"model.evaluate(test_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:06:59.421365Z","iopub.execute_input":"2026-03-05T10:06:59.421646Z","iopub.status.idle":"2026-03-05T10:07:26.136910Z","shell.execute_reply.started":"2026-03-05T10:06:59.421625Z","shell.execute_reply":"2026-03-05T10:07:26.136285Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model = tf.keras.models.load_model('/kaggle/working/best_model.h5',custom_objects={'combined_loss_penalizing_dice_and_sensitivity': combined_loss_penalizing_dice_and_sensitivity,'iou_score':iou_score,'dice_coefficient':dice_coefficient})\nbest_model.evaluate(test_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:07:26.138136Z","iopub.execute_input":"2026-03-05T10:07:26.138412Z","iopub.status.idle":"2026-03-05T10:07:56.281402Z","shell.execute_reply.started":"2026-03-05T10:07:26.138392Z","shell.execute_reply":"2026-03-05T10:07:56.280780Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# img = imread('/kaggle/working/dataset/test/img/tomo_adc026_slice_157.png')\n# img = resize(img,IMG_SIZE)\nimg = io.imread('/kaggle/working/dataset/test/img/tomo_adc026_slice_157.png')\nmask = io.imread('/kaggle/working/dataset/test/mask/tomo_adc026_slice_157.png')\nprint(f\"Initial image shape: {img.shape}\")\n\n# 2. Add the 3rd channel for RGB conversion\n# This step is necessary if the model expects a 3-channel image.\nif img.ndim == 2:\n    img = color.gray2rgb(img)\n    print(f\"Shape after converting to RGB: {img.shape}\")\n\n# 3. Resize the image with anti-aliasing\n# The `transform.resize` function handles 3-channel images automatically.\n# Use `anti_aliasing=True` for better image quality, especially when downsampling.\nresized_img = resize(\n    img,\n    IMG_SIZE,\n    anti_aliasing=True\n)\nresized_msk = resize(\n    mask,\n    IMG_SIZE,\n    anti_aliasing=True\n)\nprint(f\"resized image shape: {resized_img.shape}\")\nplt.subplot(1,2,1)\nplt.title(\"Original image\")\nplt.imshow(resized_img)\nplt.subplot(1,2,2)\nplt.title(\"Mask\")\nplt.imshow(resized_msk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:07:56.282224Z","iopub.execute_input":"2026-03-05T10:07:56.282496Z","iopub.status.idle":"2026-03-05T10:07:56.686845Z","shell.execute_reply.started":"2026-03-05T10:07:56.282475Z","shell.execute_reply":"2026-03-05T10:07:56.686044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_img = np.expand_dims(resized_img, axis=0)\npred_1 = model.predict(input_img)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:07:56.688374Z","iopub.execute_input":"2026-03-05T10:07:56.688619Z","iopub.status.idle":"2026-03-05T10:08:07.673078Z","shell.execute_reply.started":"2026-03-05T10:07:56.688602Z","shell.execute_reply":"2026-03-05T10:08:07.672427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(pred_1[0].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:08:07.673915Z","iopub.execute_input":"2026-03-05T10:08:07.674123Z","iopub.status.idle":"2026-03-05T10:08:07.678715Z","shell.execute_reply.started":"2026-03-05T10:08:07.674105Z","shell.execute_reply":"2026-03-05T10:08:07.677997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(pred_1[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:08:07.679465Z","iopub.execute_input":"2026-03-05T10:08:07.679747Z","iopub.status.idle":"2026-03-05T10:08:07.875410Z","shell.execute_reply.started":"2026-03-05T10:08:07.679722Z","shell.execute_reply":"2026-03-05T10:08:07.874705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_2 = best_model.predict(input_img)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:08:07.876221Z","iopub.execute_input":"2026-03-05T10:08:07.876540Z","iopub.status.idle":"2026-03-05T10:08:15.836129Z","shell.execute_reply.started":"2026-03-05T10:08:07.876522Z","shell.execute_reply":"2026-03-05T10:08:15.835507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(pred_2[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:08:15.837313Z","iopub.execute_input":"2026-03-05T10:08:15.837872Z","iopub.status.idle":"2026-03-05T10:08:15.996282Z","shell.execute_reply.started":"2026-03-05T10:08:15.837844Z","shell.execute_reply":"2026-03-05T10:08:15.995573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport tensorflow as tf\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\n\ndef plot_training_curves(history):\n    plt.figure(figsize=(12, 5))\n\n    # Plot Loss\n    plt.subplot(1, 2, 1)\n    plt.plot(history.history['loss'], label='Train Loss')\n    plt.plot(history.history['val_loss'], label='Val Loss')\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.grid(True)\n\n    metric_key = None\n    for k in history.history.keys():\n        if 'dice' in k.lower():\n            metric_key = k\n            break\n        elif 'accuracy' in k.lower():\n            metric_key = k\n            break\n\n    if metric_key:\n        plt.subplot(1, 2, 2)\n        plt.plot(history.history[metric_key], label=f'Train {metric_key}')\n        plt.plot(history.history[f'val_{metric_key}'], label=f'Val {metric_key}')\n        plt.title(f'Training and Validation {metric_key}')\n        plt.xlabel('Epochs')\n        plt.ylabel(metric_key)\n        plt.legend()\n        plt.grid(True)\n\n    plt.tight_layout()\n    plt.show()\n\n\ndef plot_confusion_matrix(model, test_ds, threshold=0.5):\n    y_true = []\n    y_pred = []\n\n    for images, masks in test_ds:\n        preds = model.predict(images, verbose=0)\n        preds = (preds > threshold).astype(np.uint8)  \n        y_true.append(masks.numpy().flatten())\n        y_pred.append(preds.flatten())\n\n    y_true = np.concatenate(y_true)\n    y_pred = np.concatenate(y_pred)\n\n    cm = confusion_matrix(y_true, y_pred)\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=['Background', 'Object'])\n    disp.plot(cmap=plt.cm.Blues, values_format='d')\n    plt.title(\"Confusion Matrix (Flattened Pixel-wise)\")\n    plt.show()\n\n\nplot_training_curves(history)\n\n# plot_confusion_matrix(model, test_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T10:08:15.997058Z","iopub.execute_input":"2026-03-05T10:08:15.997365Z","iopub.status.idle":"2026-03-05T10:08:16.111977Z","shell.execute_reply.started":"2026-03-05T10:08:15.997346Z","shell.execute_reply":"2026-03-05T10:08:16.111084Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# temp","metadata":{}},{"cell_type":"code","source":"# # import tensorflow as tf\n# # import os\n\n# # IMG_SIZE = (256, 256)\n# # MASK_SIZE = (256, 256)\n\n# # def load_image_and_mask(img_path, mask_path):\n# #     \"\"\"\n# #     Loads and preprocesses a single image and mask from file paths.\n# #     \"\"\"\n# #     img = tf.io.read_file(img_path)\n# #     img = tf.image.decode_png(img, channels=3)\n# #     img = tf.image.resize(img, IMG_SIZE)\n# #     img = tf.cast(img, tf.float32) / 255.0\n\n# #     mask = tf.io.read_file(mask_path)\n# #     mask = tf.image.decode_png(mask, channels=1)\n# #     mask = tf.image.resize(mask, MASK_SIZE)\n# #     mask = tf.cast(mask, tf.float32) / 255.0\n\n# #     return img, mask\n\n# # def augment(img, mask):\n# #     \"\"\"\n# #     Applies on-the-fly data augmentation to the image and mask.\n# #     \"\"\"\n# #     # Random horizontal flip\n# #     if tf.random.uniform(()) > 0.5:\n# #         img = tf.image.flip_left_right(img)\n# #         mask = tf.image.flip_left_right(mask)\n\n# #     # Random vertical flip\n# #     if tf.random.uniform(()) > 0.5:\n# #         img = tf.image.flip_up_down(img)\n# #         mask = tf.image.flip_up_down(mask)\n    \n# #     # Random rotation (90, 180, 270 degrees)\n# #     k = tf.random.uniform(shape=[], minval=0, maxval=4, dtype=tf.int32)\n# #     img = tf.image.rot90(img, k=k)\n# #     mask = tf.image.rot90(mask, k=k)\n\n# #     # Note: Color augmentations should only be applied to the image\n# #     img = tf.image.random_brightness(img, max_delta=0.1)\n# #     img = tf.image.random_contrast(img, lower=0.9, upper=1.1)\n\n# #     return img, mask\n\n# # def create_heatmap_dataset(image_dir, mask_dir, batch_size=16, shuffle=True):\n# #     \"\"\"\n# #     Creates a tf.data.Dataset pipeline for images and masks with augmentation.\n# #     \"\"\"\n# #     image_files = sorted([f for f in os.listdir(image_dir) if f.endswith('.png')])\n# #     mask_files = sorted([f for f in os.listdir(mask_dir) if f.endswith('.png')])\n\n# #     image_paths = [os.path.join(image_dir, f) for f in image_files]\n# #     mask_paths = [os.path.join(mask_dir, f) for f in mask_files]\n    \n# #     # Use from_tensor_slices for efficiency\n# #     dataset = tf.data.Dataset.from_tensor_slices((image_paths, mask_paths))\n    \n# #     if shuffle:\n# #         dataset = dataset.shuffle(buffer_size=1024)\n        \n# #     # Map the loading and parsing function\n# #     dataset = dataset.map(\n# #         lambda x, y: tf.py_function(load_image_and_mask, [x, y], (tf.float32, tf.float32)),\n# #         num_parallel_calls=tf.data.AUTOTUNE\n# #     )\n    \n# #     # Ensure correct tensor shapes after the py_function call\n# #     dataset = dataset.map(\n# #         lambda x, y: (tf.ensure_shape(x, (256, 256, 3)), tf.ensure_shape(y, (256, 256, 1)))\n# #     )\n\n# #     # Apply the augmentation function\n# #     dataset = dataset.map(\n# #         augment, \n# #         num_parallel_calls=tf.data.AUTOTUNE\n# #     )\n\n# #     # Batch and prefetch the dataset for training\n# #     return dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)","metadata":{"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_coords = pred[0]  # Extract from batch\n# print(\"Predicted coordinates (x, y):\", pred_coords)\n\n# plt.imshow(resized_img.astype('uint8'))  # Or .numpy() if using TensorFlow tensor\n# plt.scatter(pred_coords[0], pred_coords[1], c='red', marker='x')\n# plt.title(\"Predicted Motor Location\")\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-19T04:20:02.360079Z","iopub.execute_input":"2025-09-19T04:20:02.360365Z","iopub.status.idle":"2025-09-19T04:20:02.514717Z","shell.execute_reply.started":"2025-09-19T04:20:02.360348Z","shell.execute_reply":"2025-09-19T04:20:02.513912Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def generate_3d_gaussian_heatmap(shape, center, radius):\n#     \"\"\"\n#     Generates a 3D Gaussian heatmap.\n\n#     Parameters:\n#         shape  : Tuple[int, int, int] - Shape of the 3D volume (depth, height, width)\n#         center : Tuple[int, int, int] - (z, y, x) coordinates of the center\n#         radius : float                - Radius (standard deviation) of the Gaussian\n\n#     Returns:\n#         heatmap: 3D NumPy array of shape `shape`\n#     \"\"\"\n#     z = np.arange(0, shape[0])\n#     y = np.arange(0, shape[1])\n#     x = np.arange(0, shape[2])\n#     zz, yy, xx = np.meshgrid(z, y, x, indexing='ij')\n\n#     z0, y0, x0 = center\n\n#     heatmap = np.exp(-((xx - x0) ** 2 + (yy - y0) ** 2 + (zz - z0) ** 2) / (2 * radius ** 2))\n\n#     # Normalize to [0, 1]\n#     heatmap /= np.max(heatmap)\n#     return heatmap\n\n# # Example usage\n# volume_shape = (500,  924,956)     # (depth, height, width)\n# center_point = (235,  403,137)       # Center of the blob\n# radius = 10\n\n# heatmap_3d = generate_3d_gaussian_heatmap(volume_shape, center_point, radius)\n\n# # Visualize a few central slices\n# import matplotlib.pyplot as plt\n\n# fig, axs = plt.subplots(1, 3, figsize=(15, 5))\n# axs[0].imshow(heatmap_3d[center_point[0], :, :], cmap='hot')  # Z slice\n# axs[0].set_title(\"Axial Slice (Z)\")\n# axs[1].imshow(heatmap_3d[:, center_point[1], :], cmap='hot')  # Y slice\n# axs[1].set_title(\"Coronal Slice (Y)\")\n# axs[2].imshow(heatmap_3d[:, :, center_point[2]], cmap='hot')  # X slice\n# axs[2].set_title(\"Sagittal Slice (X)\")\n# plt.show()\n\n# def pad_to_shape(volume, target_shape,value):\n#     \"\"\"\n#     Pads a 3D volume to the target shape using constant padding (0).\n#     \"\"\"\n#     pad_width = []\n#     for i in range(3):  # For z, y, x\n#         total_pad = target_shape[i] - volume.shape[i]\n#         pad_before = total_pad // 2\n#         pad_after = total_pad - pad_before\n#         pad_width.append((pad_before, pad_after))\n    \n#     return np.pad(volume, pad_width, mode='constant', constant_values=value)\n\n\n# def get_volume_and_mask(sample, folder_path=TRAIN_DIR, trust=120, radius=60):\n#     tomo_id = sample['tomo_id']\n#     files = sorted(os.listdir(os.path.join(folder_path, tomo_id)))\n#     z = int(sample['Motor axis 0'])\n#     y = int(sample['Motor axis 1'])\n#     x = int(sample['Motor axis 2'])\n#     slice_dir = os.path.join(folder_path, tomo_id)\n#     volume = np.stack([\n#         img_as_float(imread(os.path.join(slice_dir, f)))\n#         for f in files\n#     ])\n#     mask = generate_3d_gaussian_heatmap(volume.shape, (z, y, x), radius)\n\n#     # Extract patch\n#     z1, z2 = max(0, z - trust), min(volume.shape[0], z + trust)\n#     y1, y2 = max(0, y - trust), min(volume.shape[1], y + trust)\n#     x1, x2 = max(0, x - trust), min(volume.shape[2], x + trust)\n\n#     patch_vol = volume[z1:z2, y1:y2, x1:x2]\n#     patch_mask = mask[z1:z2, y1:y2, x1:x2]\n\n#     # Pad if necessary\n#     target_shape = (2 * trust, 2 * trust, 2 * trust)\n#     patch_vol = pad_to_shape(patch_vol, target_shape,1)\n#     patch_mask = pad_to_shape(patch_mask, target_shape,0)\n\n#     return patch_vol.astype(np.float32), patch_mask.astype(np.float32)\n\n\n# vol, mask = get_volume_and_mask(sample,TRAIN_DIR)\n# print(f\"vol shape: {vol.shape}, mask shape: {mask.shape}\")\n\n# import gc\n# # import segmentation_models_3D as sm\n# import keras\n# import keras.backend as K\n# import tensorflow as tf\n\n# CHANNELS = 3\n\n\n\n# def datagenerator(samples,batch_size=8,DATASET_DIR=TRAIN_DIR,TRUST=120,REDIUS=40):\n#     while True:\n#         x_batch = []\n#         y_batch = []\n\n#         # Collect one full batch\n#         for _ in range(batch_size):\n#             sample = samples.sample(n=1).iloc[0]\n#             volume, mask = get_volume_and_mask(sample)\n\n#             # Add channel dimensions: volume → (D, H, W, 3), mask → (D, H, W, 1)\n#             volume = np.repeat(volume[..., np.newaxis], CHANNELS, axis=-1)\n#             mask = mask[..., np.newaxis]\n\n            \n#             x_batch.append(volume)\n#             y_batch.append(mask)\n\n#         # Stack into batch: (B, D, H, W, C)\n#         x_batch = np.stack(x_batch, axis=0)\n#         y_batch = np.stack(y_batch, axis=0)\n\n#         yield x_batch, y_batch\n#         del volume, mask, x_batch, y_batch\n#         gc.collect()\n\n# train_labels_2 = train_labels[train_labels['Number of motors']>0]\n# train_labels_2.describe()\n\n# print(f\"Total no of samples: {len(train_labels_2['tomo_id'])}\")\n\n# train_split = int(len(train_labels_2['tomo_id'])*0.70)\n# test_split = int(len(train_labels_2)*0.15)\n# val_split = int(len(train_labels_2)*0.15)\n# train_samples = train_labels_2.iloc[:train_split]\n# test_samples = train_labels_2.iloc[train_split:train_split+test_split]\n# val_samples = train_labels_2.iloc[train_split+test_split:]\n# print(f\"no of train samples: {len(train_samples['tomo_id'])}\")\n# print(f\"no of test samples: {len(test_samples['tomo_id'])}\")\n# print(f\"no of val samples: {len(val_samples['tomo_id'])}\")\n\n# train_gen = datagenerator(train_samples)\n# test_gen = datagenerator(test_samples)\n# val_gen = datagenerator(val_samples)\n\n# x_batch, y_batch = next(train_gen)\n# print(f\"x batch shape: {x_batch.shape}, y_batch shape: {y_batch.shape}\")\n\n# from tensorflow.keras.layers import Input, Conv3D, MaxPooling3D, UpSampling3D, Dropout, concatenate\n# from tensorflow.keras.models import Model\n# import tensorflow as tf\n\n# def build_unet(input_shape=(240, 240, 240, 3),num_classes=1, ker_init='he_normal', dropout=0.3):\n#     inputs = Input(input_shape)\n\n#     # Encoder\n#     conv1 = Conv3D(32, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(inputs)\n#     conv1 = Conv3D(32, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv1)\n#     pool1 = MaxPooling3D(pool_size=(2, 2, 2))(conv1)\n\n#     conv2 = Conv3D(64, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(pool1)\n#     conv2 = Conv3D(64, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv2)\n#     pool2 = MaxPooling3D(pool_size=(2, 2, 2))(conv2)\n\n#     conv3 = Conv3D(128, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(pool2)\n#     conv3 = Conv3D(128, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv3)\n#     pool3 = MaxPooling3D(pool_size=(2, 2, 2))(conv3)\n\n#     conv4 = Conv3D(256, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(pool3)\n#     conv4 = Conv3D(256, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv4)\n#     pool4 = MaxPooling3D(pool_size=(2, 2, 2))(conv4)\n\n#     # Bottleneck\n#     conv5 = Conv3D(512, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(pool4)\n#     conv5 = Conv3D(512, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv5)\n#     drop5 = Dropout(dropout)(conv5)\n\n#     # Decoder\n#     up6 = UpSampling3D(size=(2, 2, 2))(drop5)\n#     up6 = Conv3D(256, (2, 2, 2), activation='relu', padding='same', kernel_initializer=ker_init)(up6)\n#     merge6 = concatenate([conv4, up6])\n#     conv6 = Conv3D(256, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(merge6)\n#     conv6 = Conv3D(256, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv6)\n\n#     up7 = UpSampling3D(size=(2, 2, 2))(conv6)\n#     up7 = Conv3D(128, (2, 2, 2), activation='relu', padding='same', kernel_initializer=ker_init)(up7)\n#     merge7 = concatenate([conv3, up7])\n#     conv7 = Conv3D(128, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(merge7)\n#     conv7 = Conv3D(128, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv7)\n\n#     up8 = UpSampling3D(size=(2, 2, 2))(conv7)\n#     up8 = Conv3D(64, (2, 2, 2), activation='relu', padding='same', kernel_initializer=ker_init)(up8)\n#     merge8 = concatenate([conv2, up8])\n#     conv8 = Conv3D(64, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(merge8)\n#     conv8 = Conv3D(64, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv8)\n\n#     up9 = UpSampling3D(size=(2, 2, 2))(conv8)\n#     up9 = Conv3D(32, (2, 2, 2), activation='relu', padding='same', kernel_initializer=ker_init)(up9)\n#     merge9 = concatenate([conv1, up9])\n#     conv9 = Conv3D(32, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(merge9)\n#     conv9 = Conv3D(32, (3, 3, 3), activation='relu', padding='same', kernel_initializer=ker_init)(conv9)\n\n#     # Output layer\n#     conv10 = Conv3D(1, (1, 1, 1), activation='sigmoid')(conv9)\n\n#     model = Model(inputs=inputs, outputs=conv10)\n#     return model\n\n# from keras.saving import register_keras_serializable\n# @register_keras_serializable()\n# def dice_coef(y_true, y_pred, smooth=1.0):\n#     class_num = 1\n#     for i in range(class_num):\n#         y_true_f = K.flatten(y_true[:,:,:,i])\n#         y_pred_f = K.flatten(y_pred[:,:,:,i])\n#         intersection = K.sum(y_true_f * y_pred_f)\n#         loss = ((2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth))\n#    #     K.print_tensor(loss, message='loss value for class {} : '.format(SEGMENT_CLASSES[i]))\n#         if i == 0:\n#             total_loss = loss\n#         else:\n#             total_loss = total_loss + loss\n#     total_loss = total_loss / class_num\n# #    K.print_tensor(total_loss, message=' total dice coef: ')\n#     return total_loss\n\n\n \n# # define per class evaluation of dice coef\n# # inspired by https://github.com/keras-team/keras/issues/9395\n# @register_keras_serializable()\n# def dice_coef_necrotic(y_true, y_pred, epsilon=1e-6):\n#     intersection = K.sum(K.abs(y_true[:,:,:,1] * y_pred[:,:,:,1]))\n#     return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,1])) + K.sum(K.square(y_pred[:,:,:,1])) + epsilon)\n# @register_keras_serializable()\n# def dice_coef_edema(y_true, y_pred, epsilon=1e-6):\n#     intersection = K.sum(K.abs(y_true[:,:,:,2] * y_pred[:,:,:,2]))\n#     return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,2])) + K.sum(K.square(y_pred[:,:,:,2])) + epsilon)\n# @register_keras_serializable()\n# def dice_coef_enhancing(y_true, y_pred, epsilon=1e-6):\n#     intersection = K.sum(K.abs(y_true[:,:,:,3] * y_pred[:,:,:,3]))\n#     return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,3])) + K.sum(K.square(y_pred[:,:,:,3])) + epsilon)\n\n\n\n# # Computing Precision \n# @register_keras_serializable()\n# def precision(y_true, y_pred):\n#         true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))\n#         predicted_positives = K.sum(K.round(K.clip(y_pred, 0, 1)))\n#         precision = true_positives / (predicted_positives + K.epsilon())\n#         return precision\n\n    \n# # Computing Sensitivity   \n# @register_keras_serializable()\n# def sensitivity(y_true, y_pred):\n#     true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))\n#     possible_positives = K.sum(K.round(K.clip(y_true, 0, 1)))\n#     return true_positives / (possible_positives + K.epsilon())\n\n\n# # Computing Specificity\n# @register_keras_serializable()\n# def specificity(y_true, y_pred):\n#     true_negatives = K.sum(K.round(K.clip((1-y_true) * (1-y_pred), 0, 1)))\n#     possible_negatives = K.sum(K.round(K.clip(1-y_true, 0, 1)))\n#     return true_negatives / (possible_negatives + K.epsilon())\n\n\n# @register_keras_serializable()\n# def iou_3d(y_true, y_pred, threshold=0.5, smooth=1e-6):\n#     \"\"\"\n#     Calculates 3D IoU for batches of volumetric predictions.\n    \n#     Args:\n#         y_true: Ground truth tensor of shape (B, D, H, W, 1)\n#         y_pred: Predicted tensor of shape (B, D, H, W, 1)\n#         threshold: Threshold to binarize predictions\n#         smooth: Smoothing factor to avoid division by zero\n\n#     Returns:\n#         IoU score (scalar tensor)\n#     \"\"\"\n#     # Binarize prediction\n#     y_pred_bin = tf.cast(y_pred > threshold, tf.float32)\n#     y_true_bin = tf.cast(y_true > threshold, tf.float32)\n\n#     # Flatten\n#     y_pred_f = tf.reshape(y_pred_bin, [tf.shape(y_pred_bin)[0], -1])\n#     y_true_f = tf.reshape(y_true_bin, [tf.shape(y_true_bin)[0], -1])\n\n#     intersection = tf.reduce_sum(y_pred_f * y_true_f, axis=1)\n#     union = tf.reduce_sum(y_pred_f + y_true_f, axis=1) - intersection\n\n#     iou = (intersection + smooth) / (union + smooth)\n#     return tf.reduce_mean(iou)\n\n# import tensorflow.keras.backend as K\n# @register_keras_serializable()\n# def tversky_loss(y_true, y_pred, alpha=0.3, beta=0.7, smooth=1e-6):\n#     y_true_f = K.flatten(y_true)\n#     y_pred_f = K.flatten(y_pred)\n\n#     TP = K.sum(y_true_f * y_pred_f)\n#     FP = K.sum((1 - y_true_f) * y_pred_f)\n#     FN = K.sum(y_true_f * (1 - y_pred_f))\n\n#     return 1 - (TP + smooth) / (TP + alpha * FP + beta * FN + smooth)\n# @register_keras_serializable()\n# def dice_coef(y_true, y_pred, smooth=1e-6):\n#     y_true_f = K.flatten(y_true)\n#     y_pred_f = K.flatten(y_pred)\n#     intersection = K.sum(y_true_f * y_pred_f)\n#     return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n# @register_keras_serializable()\n# def combined_loss_penalizing_dice_and_sensitivity(y_true, y_pred, alpha=0.3, beta=0.7, lambda_dice=1.0):\n#     # Global Tversky loss to focus on sensitivity\n#     tversky = tversky_loss(y_true, y_pred, alpha=alpha, beta=beta)\n\n#     # Dice loss for whole mask\n#     dice_1 = dice_coef_edema(y_true, y_pred)\n#     dice_2 = dice_coef(y_true, y_pred)\n#     dice_loss_1 = 1 - dice_1\n#     dice_loss_2 = 1 - dice_2\n\n#     # Combine\n#     return tversky + lambda_dice * (dice_loss_1 +dice_loss_2)\n\n\n\n# # Final Dice Coefficient for Metrics\n# @register_keras_serializable()\n# def dice_coef_metric(y_true, y_pred, smooth=1.0):\n#     y_true_f = K.flatten(y_true)\n#     y_pred_f = K.flatten(y_pred)\n#     intersection = K.sum(y_true_f * y_pred_f)\n#     union = K.sum(y_true_f) + K.sum(y_pred_f)\n#     return (2. * intersection + smooth) / (union + smooth)\n\n\n\n# model = build_unet(num_classes=1)\n# # Compile\n# model.compile(\n#     optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n#     loss=combined_loss_penalizing_dice_and_sensitivity,\n#     metrics=[\n#         dice_coef_metric,\n#         precision,\n#         sensitivity,\n#         specificity,\n#         dice_coef_necrotic,\n#         dice_coef_edema,\n#         dice_coef_enhancing,\n#         iou_3d\n#     ]\n# )\n\n\n# model.summary()\n\n# from tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping\n\n# checkpoint_cb = ModelCheckpoint(\n#     \"best_model.keras\",                     # File to save to\n#     monitor=\"val_dice_bce_loss\",                 # Metric to monitor\n#     mode=\"min\",                         # Minimize the loss\n#     save_best_only=True,               # Only save when it's the best so far\n#     save_weights_only=False,           # Save full model\n#     verbose=1\n# )\n# checkpoint_cb_2 = ModelCheckpoint(\n#     \"best_model_dice_coef_edema.keras\",                     # File to save to\n#     monitor=\"val_dice_coef_edema\",                 # Metric to monitor\n#     mode=\"max\",                         # Minimize the loss\n#     save_best_only=True,               # Only save when it's the best so far\n#     save_weights_only=False,           # Save full model\n#     verbose=1\n# )\n# # EarlyStopping to stop training if no improvement in 5 epochs\n# early_stopping_cb = EarlyStopping(\n#     monitor=\"val_dice_coef_edema\",\n#     mode=\"min\",\n#     patience=5,\n#     restore_best_weights=True,   # Optional: restores weights from best epoch\n#     verbose=1\n# )\n\n# history = model.fit(\n#     train_gen,\n#     validation_data=val_gen,\n#     epochs=15,\n#     callbacks=[checkpoint_cb_2,early_stopping_cb]\n# )\n#     # steps_per_epoch=350,\n#     # validation_steps=100,\n\n# steps = 100  # or any value you want\n# results_2 = model.evaluate(test_gen, verbose=1)\n\n# # Show metrics\n# for name, value in zip(model.metrics_names, results_2):\n#     print(f\"{name}: {value:.4f}\")\n\n# # results_2 = model.evaluate(extra_test_gen, steps=steps, verbose=1)\n\n# # Debug output\n# print(\"Returned metrics:\", results_2)\n# print(\"Metric names:\", model.metrics_names)\n# print(f\"Length of results: {len(results_2)}, Length of metric names: {len(model.metrics_names)}\")\n\n# # Show metrics safely\n# if len(results_2) == len(model.metrics_names):\n#     for name, value in zip(model.metrics_names, results_2):\n#         print(f\"{name}: {value:.4f}\")\n# else:\n#     print(\"Mismatch between metrics_names and evaluation results.\")\n#     for idx, value in enumerate(results_2):\n#         print(f\"Metric {idx}: {value:.4f}\")","metadata":{"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null}]}