{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\n\n# # Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# # Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\n# import kagglehub\n# # kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5"},"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\nimport matplotlib.pyplot as plt\nfrom skimage.io import imread, imsave\nfrom skimage.transform import resize\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\nfrom tqdm import tqdm\nimport os\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2026-07-31T04:55:15.100292Z","iopub.execute_input":"2026-07-31T04:55:15.100651Z","iopub.status.idle":"2026-07-31T04:55:30.576635Z","shell.execute_reply.started":"2026-07-31T04:55:15.100612Z","shell.execute_reply":"2026-07-31T04:55:30.576041Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Import Dataset","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/competitions/byu-locating-bacterial-flagellar-motors-2025/train'\n\ntrain_labels = pd.read_csv('/kaggle/input/competitions/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:55:30.577690Z","iopub.execute_input":"2026-07-31T04:55:30.578284Z","iopub.status.idle":"2026-07-31T04:55:30.596707Z","shell.execute_reply.started":"2026-07-31T04:55:30.578259Z","shell.execute_reply":"2026-07-31T04:55:30.596185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_labels.iloc[1]\nsample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:55:30.597617Z","iopub.execute_input":"2026-07-31T04:55:30.598438Z","iopub.status.idle":"2026-07-31T04:55:30.611419Z","shell.execute_reply.started":"2026-07-31T04:55:30.598415Z","shell.execute_reply":"2026-07-31T04:55:30.610500Z"}},"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-07-31T04:55:30.613595Z","iopub.execute_input":"2026-07-31T04:55:30.613839Z","iopub.status.idle":"2026-07-31T04:55:30.643920Z","shell.execute_reply.started":"2026-07-31T04:55:30.613818Z","shell.execute_reply":"2026-07-31T04:55:30.643383Z"}},"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-07-31T04:55:30.644878Z","iopub.execute_input":"2026-07-31T04:55:30.645467Z","iopub.status.idle":"2026-07-31T04:55:31.113639Z","shell.execute_reply.started":"2026-07-31T04:55:30.645433Z","shell.execute_reply":"2026-07-31T04:55:31.112923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(img.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:55:31.114490Z","iopub.execute_input":"2026-07-31T04:55:31.114886Z","iopub.status.idle":"2026-07-31T04:55:31.124290Z","shell.execute_reply.started":"2026-07-31T04:55:31.114857Z","shell.execute_reply":"2026-07-31T04:55:31.123503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_mask(img,center,sigma=3.0,quantile_=0.70,length=25):\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=40)\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-07-31T04:55:31.125118Z","iopub.execute_input":"2026-07-31T04:55:31.125736Z","iopub.status.idle":"2026-07-31T04:55:32.405724Z","shell.execute_reply.started":"2026-07-31T04:55:31.125713Z","shell.execute_reply":"2026-07-31T04:55:32.405013Z"}},"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-07-31T04:55:32.406712Z","iopub.execute_input":"2026-07-31T04:55:32.407058Z","iopub.status.idle":"2026-07-31T04:55:32.448900Z","shell.execute_reply.started":"2026-07-31T04:55:32.407034Z","shell.execute_reply":"2026-07-31T04:55:32.448339Z"}},"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-07-31T04:55:32.449832Z","iopub.execute_input":"2026-07-31T04:55:32.450177Z","iopub.status.idle":"2026-07-31T04:55:32.556001Z","shell.execute_reply.started":"2026-07-31T04:55:32.450153Z","shell.execute_reply":"2026-07-31T04:55:32.555342Z"}},"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-07-31T04:55:32.558771Z","iopub.execute_input":"2026-07-31T04:55:32.559009Z","iopub.status.idle":"2026-07-31T04:55:32.564808Z","shell.execute_reply.started":"2026-07-31T04:55:32.558989Z","shell.execute_reply":"2026-07-31T04:55:32.564053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXTRA_DATASET_DIR = '/kaggle/input/datasets/brendanartley/cryoet-flagellar-motors-dataset/jpgs'\nextra_dataset_labels = pd.read_csv('/kaggle/input/datasets/brendanartley/cryoet-flagellar-motors-dataset/labels.csv')\nextra_dataset_labels_new = pd.read_csv('/kaggle/input/datasets/brendanartley/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-07-31T04:55:32.565884Z","iopub.execute_input":"2026-07-31T04:55:32.566264Z","iopub.status.idle":"2026-07-31T04:55:32.667182Z","shell.execute_reply.started":"2026-07-31T04:55:32.566221Z","shell.execute_reply":"2026-07-31T04:55:32.666556Z"}},"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-07-31T04:55:32.668090Z","iopub.execute_input":"2026-07-31T04:55:32.668349Z","iopub.status.idle":"2026-07-31T04:55:32.703887Z","shell.execute_reply.started":"2026-07-31T04:55:32.668327Z","shell.execute_reply":"2026-07-31T04:55:32.703301Z"}},"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=25):\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=40):\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    # # 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    mask_dir = os.path.join(output_base_dir, 'mask')\n    os.makedirs(mask_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    if slice_idx -15 > 0:\n        files = sorted(os.listdir(os.path.join(dataset_dir, tomo_id)))\n        file_name = files[int(slice_idx-15)]\n        path_2 = os.path.join(dataset_dir,tomo_id,file_name)\n        image2 = imread(path_2)\n        img_path2 = os.path.join(output_base_dir, 'img', f\"{ tomo_id}_slice_{slice_idx-15}.png\")\n        mask_path2 = os.path.join(output_base_dir, 'mask',  f\"{tomo_id}_slice_{slice_idx-15}.png\")\n        imsave(img_path2, image2)\n        # print(image_shape[0])\n        imsave(mask_path2,np.zeros(image_shape, dtype=np.uint8))\n\n    elif slice_idx +15 < len(os.listdir(os.path.join(dataset_dir, tomo_id))):\n        files = sorted(os.listdir(os.path.join(dataset_dir, tomo_id)))\n        file_name = files[int(slice_idx+15)]\n        path_2 = os.path.join(dataset_dir,tomo_id,file_name)\n        image2 = imread(path_2)\n        img_path2 = os.path.join(output_base_dir, 'img', f\"{ tomo_id}_slice_{slice_idx+15}.png\")\n        mask_path2 = os.path.join(output_base_dir, 'mask',  f\"{tomo_id}_slice_{slice_idx+15}.png\")\n        imsave(img_path2, image2)\n        imsave(mask_path2,np.zeros(image_shape, dtype=np.uint8))\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    # heatmap = generate_gaussian_heatmap(image_shape,(int(center_x), int(center_y)),radius)\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-07-31T04:55:32.704792Z","iopub.execute_input":"2026-07-31T04:55:32.705089Z","iopub.status.idle":"2026-07-31T04:55:32.722422Z","shell.execute_reply.started":"2026-07-31T04:55:32.705042Z","shell.execute_reply":"2026-07-31T04:55:32.721582Z"}},"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-07-31T04:55:32.723780Z","iopub.execute_input":"2026-07-31T04:55:32.724055Z","iopub.status.idle":"2026-07-31T04:55:32.740060Z","shell.execute_reply.started":"2026-07-31T04:55:32.724019Z","shell.execute_reply":"2026-07-31T04:55:32.739408Z"}},"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-07-31T04:55:32.740977Z","iopub.execute_input":"2026-07-31T04:55:32.741282Z","iopub.status.idle":"2026-07-31T04:56:41.210811Z","shell.execute_reply.started":"2026-07-31T04:55:32.741263Z","shell.execute_reply":"2026-07-31T04:56:41.209660Z"}},"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-07-31T04:56:41.211855Z","iopub.execute_input":"2026-07-31T04:56:41.212152Z","iopub.status.idle":"2026-07-31T04:56:55.668493Z","shell.execute_reply.started":"2026-07-31T04:56:41.212124Z","shell.execute_reply":"2026-07-31T04:56:55.667568Z"}},"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-07-31T04:56:55.669525Z","iopub.execute_input":"2026-07-31T04:56:55.669813Z","iopub.status.idle":"2026-07-31T04:57:10.971095Z","shell.execute_reply.started":"2026-07-31T04:56:55.669779Z","shell.execute_reply":"2026-07-31T04:57:10.970502Z"}},"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-07-31T04:57:10.972036Z","iopub.execute_input":"2026-07-31T04:57:10.972402Z","iopub.status.idle":"2026-07-31T04:57:10.978150Z","shell.execute_reply.started":"2026-07-31T04:57:10.972378Z","shell.execute_reply":"2026-07-31T04:57:10.977314Z"}},"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-07-31T04:57:10.978997Z","iopub.execute_input":"2026-07-31T04:57:10.979236Z","iopub.status.idle":"2026-07-31T04:57:10.998057Z","shell.execute_reply.started":"2026-07-31T04:57:10.979216Z","shell.execute_reply":"2026-07-31T04:57:10.997437Z"}},"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-07-31T04:57:10.999218Z","iopub.execute_input":"2026-07-31T04:57:10.999503Z","iopub.status.idle":"2026-07-31T04:57:11.019178Z","shell.execute_reply.started":"2026-07-31T04:57:10.999477Z","shell.execute_reply":"2026-07-31T04:57:11.018425Z"}},"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-07-31T04:57:11.020065Z","iopub.execute_input":"2026-07-31T04:57:11.020329Z","iopub.status.idle":"2026-07-31T04:57:11.040419Z","shell.execute_reply.started":"2026-07-31T04:57:11.020309Z","shell.execute_reply":"2026-07-31T04:57:11.039755Z"}},"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-07-31T04:57:11.041260Z","iopub.execute_input":"2026-07-31T04:57:11.041527Z","iopub.status.idle":"2026-07-31T04:57:11.055532Z","shell.execute_reply.started":"2026-07-31T04:57:11.041493Z","shell.execute_reply":"2026-07-31T04:57:11.054745Z"}},"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-07-31T04:57:11.056583Z","iopub.execute_input":"2026-07-31T04:57:11.057285Z","iopub.status.idle":"2026-07-31T04:58:59.326508Z","shell.execute_reply.started":"2026-07-31T04:57:11.057262Z","shell.execute_reply":"2026-07-31T04:58:59.325885Z"}},"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-07-31T04:58:59.327490Z","iopub.execute_input":"2026-07-31T04:58:59.327848Z","iopub.status.idle":"2026-07-31T04:59:12.526204Z","shell.execute_reply.started":"2026-07-31T04:58:59.327824Z","shell.execute_reply":"2026-07-31T04:59:12.525004Z"}},"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-07-31T04:59:12.527377Z","iopub.execute_input":"2026-07-31T04:59:12.527667Z","iopub.status.idle":"2026-07-31T04:59:26.804993Z","shell.execute_reply.started":"2026-07-31T04:59:12.527644Z","shell.execute_reply":"2026-07-31T04:59:26.804402Z"}},"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-07-31T04:59:26.806027Z","iopub.execute_input":"2026-07-31T04:59:26.806429Z","iopub.status.idle":"2026-07-31T04:59:26.815118Z","shell.execute_reply.started":"2026-07-31T04:59:26.806387Z","shell.execute_reply":"2026-07-31T04:59:26.814159Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset & DataLoader (PyTorch)","metadata":{}},{"cell_type":"code","source":"!uv pip install -q albumentations segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:26.816177Z","iopub.execute_input":"2026-07-31T04:59:26.816483Z","iopub.status.idle":"2026-07-31T04:59:28.117497Z","shell.execute_reply.started":"2026-07-31T04:59:26.816461Z","shell.execute_reply":"2026-07-31T04:59:28.116475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch, subprocess\nprint(\"torch version:\", torch.__version__)\nprint(\"cuda available:\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"device:\", torch.cuda.get_device_name(0))\n    print(\"capability:\", torch.cuda.get_device_capability(0))\nprint(subprocess.run([\"nvidia-smi\"], capture_output=True, text=True).stdout)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:28.121959Z","iopub.execute_input":"2026-07-31T04:59:28.122366Z","iopub.status.idle":"2026-07-31T04:59:33.481559Z","shell.execute_reply.started":"2026-07-31T04:59:28.122317Z","shell.execute_reply":"2026-07-31T04:59:33.479671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python --version","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:33.483264Z","iopub.execute_input":"2026-07-31T04:59:33.483573Z","iopub.status.idle":"2026-07-31T04:59:33.634411Z","shell.execute_reply.started":"2026-07-31T04:59:33.483539Z","shell.execute_reply":"2026-07-31T04:59:33.633430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport cv2\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nIMG_SIZE = 256\nBATCH_SIZE = 16\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:33.636110Z","iopub.execute_input":"2026-07-31T04:59:33.636401Z","iopub.status.idle":"2026-07-31T04:59:34.952596Z","shell.execute_reply.started":"2026-07-31T04:59:33.636374Z","shell.execute_reply":"2026-07-31T04:59:34.951945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FlagellarDataset(Dataset):\n    \"\"\"\n    Reads matching image/mask png pairs from `image_dir` / `mask_dir`\n    (same layout produced by flagellar_dataset() above: <base>/img and <base>/mask).\n    \"\"\"\n    def __init__(self, image_dir, mask_dir, transform=None, img_size=IMG_SIZE):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.image_files = sorted([f for f in os.listdir(image_dir) if f.endswith('.png')])\n        self.mask_files = sorted([f for f in os.listdir(mask_dir) if f.endswith('.png')])\n        assert len(self.image_files) == len(self.mask_files), \\\n            f\"Mismatch: {len(self.image_files)} images vs {len(self.mask_files)} masks\"\n        self.transform = transform\n        self.img_size = img_size\n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.image_dir, self.image_files[idx])\n        mask_path = os.path.join(self.mask_dir, self.mask_files[idx])\n\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n\n        image = cv2.resize(image, (self.img_size, self.img_size))\n        mask = cv2.resize(mask, (self.img_size, self.img_size), interpolation=cv2.INTER_NEAREST)\n        mask = mask.astype(np.float32) / 255.0\n\n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented['image']      # tensor (3,H,W), normalized to [0,1]\n            mask = augmented['mask']        # tensor (H,W), float in [0,1]\n        else:\n            image = torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0\n            mask = torch.from_numpy(mask).float()\n\n        mask = mask.unsqueeze(0)  # (1, H, W) to match model output shape\n        return image, mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:34.953614Z","iopub.execute_input":"2026-07-31T04:59:34.953916Z","iopub.status.idle":"2026-07-31T04:59:34.961822Z","shell.execute_reply.started":"2026-07-31T04:59:34.953892Z","shell.execute_reply":"2026-07-31T04:59:34.961047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Augmentations mirror the TF pipeline: flips, 90-degree rotations, brightness/contrast jitter.\ntrain_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n    A.Normalize(mean=(0.0, 0.0, 0.0), std=(1.0, 1.0, 1.0), max_pixel_value=255.0),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.Normalize(mean=(0.0, 0.0, 0.0), std=(1.0, 1.0, 1.0), max_pixel_value=255.0),\n    ToTensorV2(),\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:34.962764Z","iopub.execute_input":"2026-07-31T04:59:34.962966Z","iopub.status.idle":"2026-07-31T04:59:34.986672Z","shell.execute_reply.started":"2026-07-31T04:59:34.962942Z","shell.execute_reply":"2026-07-31T04:59:34.985796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DATASET_DIR = '/kaggle/working/dataset/train'\nVAL_DATASET_DIR = '/kaggle/working/dataset/val'\nTEST_DATASET_DIR = '/kaggle/working/dataset/test'\ntrain_img_dir = os.path.join(TRAIN_DATASET_DIR, \"img\")\ntrain_mask_dir = os.path.join(TRAIN_DATASET_DIR, \"mask\")\nval_img_dir = os.path.join(VAL_DATASET_DIR, \"img\")\nval_mask_dir = os.path.join(VAL_DATASET_DIR, \"mask\")\ntest_img_dir = os.path.join(TEST_DATASET_DIR, \"img\")\ntest_mask_dir = os.path.join(TEST_DATASET_DIR, \"mask\")\n\nbase_train_ds = FlagellarDataset(train_img_dir, train_mask_dir, transform=train_transform)\n\n# The TF notebook used augment_multiplier=5 to expand the effective training set with\n# repeated random augmentations. ConcatDataset gives the same effect in PyTorch, since\n# each copy re-applies the (random) train_transform independently.\nAUGMENT_MULTIPLIER = 5\ntrain_ds = ConcatDataset([base_train_ds] * AUGMENT_MULTIPLIER)\n\nval_ds = FlagellarDataset(val_img_dir, val_mask_dir, transform=val_transform)\ntest_ds = FlagellarDataset(test_img_dir, test_mask_dir, transform=val_transform)\n\n# train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True)\n# val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,num_workers=1, pin_memory=True, persistent_workers=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,num_workers=1, pin_memory=True, persistent_workers=True)\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\n\nprint(f\"Train batches: {len(train_loader)} | Val batches: {len(val_loader)} | Test batches: {len(test_loader)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:34.987867Z","iopub.execute_input":"2026-07-31T04:59:34.988146Z","iopub.status.idle":"2026-07-31T04:59:35.006425Z","shell.execute_reply.started":"2026-07-31T04:59:34.988119Z","shell.execute_reply":"2026-07-31T04:59:35.005652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Quick sanity check: pull one batch and inspect shapes / value ranges\nimages, masks = next(iter(train_loader))\nprint(\"images:\", images.shape, images.dtype, images.min().item(), images.max().item())\nprint(\"masks: \", masks.shape, masks.dtype, masks.min().item(), masks.max().item())\n\nplt.figure(figsize=(8, 4))\nplt.subplot(1, 2, 1)\nplt.imshow(images[0].permute(1, 2, 0).numpy())\nplt.title(\"Augmented image\")\nplt.subplot(1, 2, 2)\nplt.imshow(masks[0, 0].numpy(), cmap=\"gray\")\nplt.title(\"Augmented mask\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:35.007465Z","iopub.execute_input":"2026-07-31T04:59:35.007750Z","iopub.status.idle":"2026-07-31T04:59:35.911947Z","shell.execute_reply.started":"2026-07-31T04:59:35.007730Z","shell.execute_reply":"2026-07-31T04:59:35.911104Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model (U-Net, EfficientNet-B0 encoder)","metadata":{}},{"cell_type":"code","source":"# from kaggle_secrets import UserSecretsClient\n# user_secrets = UserSecretsClient()\n# secret_value_0 = user_secrets.get_secret(\"HF_TOKEN\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:35.913696Z","iopub.execute_input":"2026-07-31T04:59:35.914170Z","iopub.status.idle":"2026-07-31T04:59:35.918486Z","shell.execute_reply.started":"2026-07-31T04:59:35.914125Z","shell.execute_reply":"2026-07-31T04:59:35.917572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import segmentation_models_pytorch as smp\n\n# # smp.Unet is the PyTorch equivalent of the encoder-decoder U-Net built manually in the\n# # TF notebook: an ImageNet-pretrained EfficientNet-B0 backbone as encoder, with skip\n# # connections automatically wired into a U-Net decoder. This replaces the manual\n# # skip_layers / Concatenate wiring from the TF version.\n# model = smp.Unet(\n#     encoder_name=\"efficientnet-b0\",\n#     encoder_weights=\"imagenet\",\n#     in_channels=3,\n#     classes=1,\n#     activation=None,   # keep raw logits out; sigmoid applied inside the loss/metrics below\n#                         # for numerical stability (equivalent to the TF model's sigmoid output)\n# )\n# model = model.to(device)\n\n# n_params = sum(p.numel() for p in model.parameters())\n# print(f\"Total parameters: {n_params:,}\")\nimport torch\nimport segmentation_models_pytorch as smp\n\n# 1. Define the model (same as your code)\nmodel = smp.Unet(\n    encoder_name=\"efficientnet-b0\",\n    encoder_weights=\"imagenet\",\n    in_channels=3,\n    classes=1,\n    activation=None,\n)\n\n# 2. Check for multiple GPUs and wrap\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = torch.nn.DataParallel(model)\n\n# 3. Move to device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\nn_params = sum(p.numel() for p in model.parameters())\nprint(f\"Total parameters: {n_params:,}\")   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:35.919391Z","iopub.execute_input":"2026-07-31T04:59:35.919735Z","iopub.status.idle":"2026-07-31T04:59:45.079355Z","shell.execute_reply.started":"2026-07-31T04:59:35.919703Z","shell.execute_reply":"2026-07-31T04:59:45.078601Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Losses & Metrics","metadata":{}},{"cell_type":"code","source":"# Same Dice / IoU / Tversky / combined loss used in the TF notebook, reimplemented for\n# PyTorch. Since the model outputs raw logits, each function applies sigmoid internally.\n\ndef dice_coefficient(y_pred_logits, y_true, smooth=1.0):\n    y_pred = torch.sigmoid(y_pred_logits)\n    y_pred_f = y_pred.reshape(-1)\n    y_true_f = y_true.reshape(-1)\n    intersection = (y_pred_f * y_true_f).sum()\n    return (2. * intersection + smooth) / (y_pred_f.sum() + y_true_f.sum() + smooth)\n\ndef dice_loss(y_pred_logits, y_true):\n    return 1 - dice_coefficient(y_pred_logits, y_true)\n\ndef iou_loss(y_pred_logits, y_true):\n    return 1 - iou_score(y_pred_logits,y_true)\n\ndef iou_score(y_pred_logits, y_true, smooth=1.0):\n    y_pred = torch.sigmoid(y_pred_logits)\n    y_pred_f = y_pred.reshape(-1)\n    y_true_f = y_true.reshape(-1)\n    intersection = (y_pred_f * y_true_f).sum()\n    union = y_pred_f.sum() + y_true_f.sum() - intersection\n    return (intersection + smooth) / (union + smooth)\n\ndef tversky_loss(y_pred_logits, y_true, alpha=0.3, beta=0.7, smooth=1e-6):\n    y_pred = torch.sigmoid(y_pred_logits)\n    y_pred_f = y_pred.reshape(-1)\n    y_true_f = y_true.reshape(-1)\n    TP = (y_true_f * y_pred_f).sum()\n    FP = ((1 - y_true_f) * y_pred_f).sum()\n    FN = (y_true_f * (1 - y_pred_f)).sum()\n    return 1 - (TP + smooth) / (TP + alpha * FP + beta * FN + smooth)\n\ndef combined_loss(y_pred_logits, y_true, alpha=0.3, beta=0.7, lambda_dice=1.0):\n    tversky = tversky_loss(y_pred_logits, y_true, alpha=alpha, beta=beta)\n    dice = dice_loss(y_pred_logits, y_true)\n    return tversky + lambda_dice * dice\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.080255Z","iopub.execute_input":"2026-07-31T04:59:45.080797Z","iopub.status.idle":"2026-07-31T04:59:45.090043Z","shell.execute_reply.started":"2026-07-31T04:59:45.080774Z","shell.execute_reply":"2026-07-31T04:59:45.089367Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\nCHECKPOINT_PATH = \"/kaggle/input/models/rajavignesh407/check-point/pytorch/default/1/best_model.pth\"\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\nmodel = smp.Unet(\n    encoder_name=\"efficientnet-b0\",\n    encoder_weights=None,\n    in_channels=3,\n    classes=1,\n    activation=None,\n)\n# 2. Check for multiple GPUs and wrap\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = torch.nn.DataParallel(model)\n\nstate_dict = torch.load(CHECKPOINT_PATH, map_location=device)\n\n# Strip \"module.\" prefix left over from DataParallel training\nnew_state_dict = {}\nfor k, v in state_dict.items():\n    new_key = k.replace(\"module.\", \"\", 1) if k.startswith(\"module.\") else k\n    new_state_dict[new_key] = v\n\nmodel.load_state_dict(state_dict)\n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.091193Z","iopub.execute_input":"2026-07-31T04:59:45.091528Z","iopub.status.idle":"2026-07-31T04:59:45.265686Z","shell.execute_reply.started":"2026-07-31T04:59:45.091497Z","shell.execute_reply":"2026-07-31T04:59:45.265125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\n\n# optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='min', factor=0.5, patience=2\n)\n\nEPOCHS = 10\nPATIENCE = 5          # early stopping patience, matches the TF EarlyStopping callback\nCHECKPOINT_PATH = \"best_model.pth\"\n\nbest_val_loss = float(\"inf\")\nepochs_no_improve = 0\nbest_model_wts = copy.deepcopy(model.state_dict())\n\nhistory = {\"train_loss\": [], \"val_loss\": [], \"train_dice\": [], \"val_dice\": [],\"train_iou\":[],\"val_iou\":[]}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.266754Z","iopub.execute_input":"2026-07-31T04:59:45.267085Z","iopub.status.idle":"2026-07-31T04:59:45.304981Z","shell.execute_reply.started":"2026-07-31T04:59:45.267033Z","shell.execute_reply":"2026-07-31T04:59:45.304422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !CUDA_LAUNCH_BLOCKING=1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.305840Z","iopub.execute_input":"2026-07-31T04:59:45.306443Z","iopub.status.idle":"2026-07-31T04:59:45.309997Z","shell.execute_reply.started":"2026-07-31T04:59:45.306420Z","shell.execute_reply":"2026-07-31T04:59:45.309420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch.multiprocessing as mp\n# mp.set_start_method('spawn', force=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.311165Z","iopub.execute_input":"2026-07-31T04:59:45.311476Z","iopub.status.idle":"2026-07-31T04:59:45.323103Z","shell.execute_reply.started":"2026-07-31T04:59:45.311454Z","shell.execute_reply":"2026-07-31T04:59:45.322459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.disable()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T04:59:45.324128Z","iopub.execute_input":"2026-07-31T04:59:45.324440Z","iopub.status.idle":"2026-07-31T04:59:45.336735Z","shell.execute_reply.started":"2026-07-31T04:59:45.324418Z","shell.execute_reply":"2026-07-31T04:59:45.335886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(dataloader,model,loss_fn,optimizer):\n    model.train()\n    running_loss, running_dice, running_iou, n_train = 0.0, 0.0, 0.0, 0\n    progress_bar = tqdm(dataloader,desc=\"Training\")\n    for images, masks in progress_bar:\n        images, masks = images.to(device), masks.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        detached_outputs = outputs.detach()\n        \n        loss = loss_fn(outputs, masks)\n        # loss = loss_fn(detached_outputs, masks)\n        loss.backward()\n        optimizer.step()\n\n        bs = images.size(0)\n        running_loss += loss.item() * bs\n        running_dice += dice_coefficient(detached_outputs, masks).item() * bs\n        running_iou += iou_score(detached_outputs,masks).item() * bs\n        n_train += bs\n\n        # current_loss = running_loss / n_train\n        # current_dice = running_dice / n_train\n        # current_iou = running_iou / n_train\n        train_loss = running_loss / n_train\n        train_dice = running_dice / n_train\n        train_iou = running_iou / n_train\n        progress_bar.set_postfix(\n            loss=f\"{train_loss:.4f}\",\n            dice=f\"{train_dice:.4f}\",\n            iou=f\"{train_iou:.4f}\",\n        )\n\n    return train_iou,train_dice,train_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T05:00:39.418710Z","iopub.execute_input":"2026-07-31T05:00:39.419449Z","iopub.status.idle":"2026-07-31T05:00:39.425844Z","shell.execute_reply.started":"2026-07-31T05:00:39.419420Z","shell.execute_reply":"2026-07-31T05:00:39.424922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def val(dataloader,model,loss_fn,optimizer):\n    model.eval()\n    val_loss_sum, val_dice_sum,val_iou_sum, n_val = 0.0, 0.0,0.0, 0\n    progress_bar = tqdm(dataloader,desc=\"Val\")\n    with torch.no_grad():\n        for images, masks in progress_bar:\n            images, masks = images.to(device), masks.to(device)\n            outputs = model(images)\n            detached_outputs = outputs.detach()\n            loss = loss_fn(outputs, masks)\n\n            bs = images.size(0)\n            val_loss_sum += loss.item() * bs\n            val_dice_sum += dice_coefficient(detached_outputs, masks).item() * bs\n            val_iou_sum += iou_score(detached_outputs,masks).item() * bs\n            n_val += bs\n\n            val_loss = val_loss_sum / n_val\n            val_dice = val_dice_sum / n_val\n            val_iou = val_iou_sum / n_val\n            \n            progress_bar.set_postfix(\n                loss=f\"{val_loss:.4f}\",\n                dice=f\"{val_dice:.4f}\",\n                iou=f\"{val_iou:.4f}\"\n            )\n        # after computing val_loss each epoch:\n        scheduler.step(val_loss)\n\n        return val_loss,val_dice,val_iou","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T05:00:41.054675Z","iopub.execute_input":"2026-07-31T05:00:41.055407Z","iopub.status.idle":"2026-07-31T05:00:41.062178Z","shell.execute_reply.started":"2026-07-31T05:00:41.055376Z","shell.execute_reply":"2026-07-31T05:00:41.061336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(EPOCHS):\n    # ---- train ----\n    # model.train()\n    # running_loss, running_dice, n_train = 0.0, 0.0, 0\n    # train_bar = tqdm(train_loader,desc=\"Training\",leave=True)\n    # for images, masks in train_loader:\n    #     images, masks = images.to(device), masks.to(device)\n\n    #     optimizer.zero_grad()\n    #     outputs = model(images)\n    #     loss = combined_loss(outputs, masks)\n    #     loss.backward()\n    #     optimizer.step()\n\n    #     bs = images.size(0)\n    #     running_loss += loss.item() * bs\n    #     running_dice += dice_coefficient(outputs, masks).item() * bs\n    #     n_train += bs\n\n    # train_loss = running_loss / n_train\n    # train_dice = running_dice / n_train\n    train_iou,train_dice,train_loss = train(train_loader,model,iou_loss,optimizer)\n\n    # ---- validate ----\n    # model.eval()\n    # val_loss_sum, val_dice_sum, n_val = 0.0, 0.0, 0\n    # with torch.no_grad():\n    #     for images, masks in val_loader:\n    #         images, masks = images.to(device), masks.to(device)\n    #         outputs = model(images)\n    #         loss = combined_loss(outputs, masks)\n\n    #         bs = images.size(0)\n    #         val_loss_sum += loss.item() * bs\n    #         val_dice_sum += dice_coefficient(outputs, masks).item() * bs\n    #         n_val += bs\n\n    # val_loss = val_loss_sum / n_val\n    # val_dice = val_dice_sum / n_val\n\n    val_loss,val_dice,val_iou = val(val_loader,model,iou_loss,optimizer)\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"train_dice\"].append(train_dice)\n    history[\"val_dice\"].append(val_dice)\n    history[\"train_iou\"].append(train_iou)\n    history[\"val_iou\"].append(val_iou)\n\n    print(f\"Epoch {epoch+1:02d}/{EPOCHS} | \"\n          f\"train_loss: {train_loss:.4f} train_dice: {train_dice:.4f} | \"\n          f\"val_loss: {val_loss:.4f} val_dice: {val_dice:.4f}\")\n    gc.collect()\n    torch.cuda.empty_cache()\n    # ---- checkpoint + early stopping (mirrors ModelCheckpoint + EarlyStopping in TF) ----\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        best_model_wts = copy.deepcopy(model.state_dict())\n        epochs_no_improve = 0\n        torch.save(model.state_dict(), CHECKPOINT_PATH)\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= PATIENCE:\n            print(f\"Early stopping triggered at epoch {epoch+1}\")\n            break\n\n# restore best weights, same as EarlyStopping(restore_best_weights=True)\nmodel.load_state_dict(best_model_wts)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-31T05:00:41.465669Z","iopub.execute_input":"2026-07-31T05:00:41.466548Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"model.eval()\ntest_loss_sum, test_dice_sum, test_iou_sum, n_test = 0.0, 0.0, 0.0, 0\n\nwith torch.no_grad():\n    for images, masks in test_loader:\n        images, masks = images.to(device), masks.to(device)\n        outputs = model(images)\n\n        bs = images.size(0)\n        test_loss_sum += combined_loss(outputs, masks).item() * bs\n        test_dice_sum += dice_coefficient(outputs, masks).item() * bs\n        test_iou_sum += iou_score(outputs, masks).item() * bs\n        n_test += bs\n\nprint(f\"Test loss: {test_loss_sum/n_test:.4f}\")\nprint(f\"Test dice: {test_dice_sum/n_test:.4f}\")\nprint(f\"Test IoU:  {test_iou_sum/n_test:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-30T18:44:12.646017Z","iopub.execute_input":"2026-07-30T18:44:12.646399Z","iopub.status.idle":"2026-07-30T18:44:18.130273Z","shell.execute_reply.started":"2026-07-30T18:44:12.646365Z","shell.execute_reply":"2026-07-30T18:44:18.129091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(history[\"train_loss\"], label=\"Train Loss\")\nplt.plot(history[\"val_loss\"], label=\"Val Loss\")\nplt.title(\"Training and Validation Loss\")\nplt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.legend(); plt.grid(True)\n\nplt.subplot(1, 2, 2)\nplt.plot(history[\"train_dice\"], label=\"Train Dice\")\nplt.plot(history[\"val_dice\"], label=\"Val Dice\")\nplt.title(\"Training and Validation Dice Coefficient\")\nplt.xlabel(\"Epoch\"); plt.ylabel(\"Dice\"); plt.legend(); plt.grid(True)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-30T18:44:18.131632Z","iopub.execute_input":"2026-07-30T18:44:18.131955Z","iopub.status.idle":"2026-07-30T18:44:18.407674Z","shell.execute_reply.started":"2026-07-30T18:44:18.131927Z","shell.execute_reply":"2026-07-30T18:44:18.406982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize a prediction from the test set\nmodel.eval()\nimages, masks = next(iter(test_loader))\nimages, masks = images.to(device), masks.to(device)\n\nwith torch.no_grad():\n    preds = torch.sigmoid(model(images))\n\nidx = 0\nimg_np = images[idx].cpu().permute(1, 2, 0).numpy()\nmask_np = masks[idx, 0].cpu().numpy()\npred_np = preds[idx, 0].cpu().numpy()\n\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 3, 1); plt.imshow(img_np); plt.title(\"Image\")\nplt.subplot(1, 3, 2); plt.imshow(mask_np, cmap=\"gray\"); plt.title(\"Ground Truth Mask\")\nplt.subplot(1, 3, 3); plt.imshow(pred_np, cmap=\"gray\"); plt.title(\"Prediction\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-30T18:44:18.408705Z","iopub.execute_input":"2026-07-30T18:44:18.408973Z","iopub.status.idle":"2026-07-30T18:44:20.595410Z","shell.execute_reply.started":"2026-07-30T18:44:18.408951Z","shell.execute_reply":"2026-07-30T18:44:20.594506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}