{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":52279,"databundleVersionId":5822112},{"sourceType":"competition","sourceId":34547,"databundleVersionId":3897958}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# [1] Pre-stage of preprocessing \n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nfrom dataclasses import dataclass\nimport colorama as c\n\n@dataclass\nclass DataPaths:\n    root_path: str = \"/kaggle/input/competitions/hubmap-organ-segmentation\"\n    train_images_paths: str = os.path.join(root_path, \"train_images\")\n    test_images_paths: str = os.path.join(root_path, \"test_images\")\n    train_annotations_path: str = os.path.join(root_path, \"train_annotations\")\n    sample_submission_path: str = os.path.join(root_path, \"sample_submission.csv\")\n    train_csv_path: str = os.path.join(root_path, \"train.csv\")\n    test_csv_path: str = os.path.join(root_path, \"test.csv\")\n\n\nprint(f\"{c.Fore.CYAN} -> Root path: {DataPaths.root_path}\")\nprint(f\"{c.Fore.CYAN} -> Train Images path: {DataPaths.train_images_paths}\")\nprint(f\"{c.Fore.CYAN} -> Test Images path: {DataPaths.test_images_paths}\")\nprint(f\"{c.Fore.CYAN} -> Train annotations path: {DataPaths.train_annotations_path}\")\nprint(f\"{c.Fore.CYAN} -> Train csv path: {DataPaths.train_csv_path}\")\nprint(f\"{c.Fore.CYAN} -> Test csv path: {DataPaths.test_csv_path}\")\nprint(f\"{c.Fore.CYAN} -> Sample submission path: {DataPaths.sample_submission_path}\")\nprint(f\"{c.Fore.GREEN} Succesfully initialized the path variables!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:03.839106Z","iopub.execute_input":"2026-05-19T15:27:03.839453Z","iopub.status.idle":"2026-05-19T15:27:04.118029Z","shell.execute_reply.started":"2026-05-19T15:27:03.839418Z","shell.execute_reply":"2026-05-19T15:27:04.117245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass MetaDataSource:\n    train_data: pd.DataFrame = pd.read_csv(DataPaths.train_csv_path)\n    test_data: pd.DataFrame = pd.read_csv(DataPaths.test_csv_path)\n    sample_submission_data: pd.DataFrame = pd.read_csv(DataPaths.sample_submission_path)\n\n\n# Printing that basic stats of train data\n\nprint(MetaDataSource.train_data.shape)\nprint(MetaDataSource.train_data.columns)\nMetaDataSource.train_data.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:07.253475Z","iopub.execute_input":"2026-05-19T15:27:07.253965Z","iopub.status.idle":"2026-05-19T15:27:07.555754Z","shell.execute_reply.started":"2026-05-19T15:27:07.253933Z","shell.execute_reply":"2026-05-19T15:27:07.554678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass RandomStates:\n    random_state: int = 14","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:10.849951Z","iopub.execute_input":"2026-05-19T15:27:10.850229Z","iopub.status.idle":"2026-05-19T15:27:10.855612Z","shell.execute_reply.started":"2026-05-19T15:27:10.850205Z","shell.execute_reply":"2026-05-19T15:27:10.854615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Column explaination (for self use)\n\n# - id: A id for each samples, useful: For loading images\n# - organ: Showing what type of organ it is . useful: This is core for multi modeling\n# - img_height: Height of the image. useful: maybe\n# - img_width: Width of the image. useful: maybe\n# - pixel_size: The pixel size in micrometers. useful: Crucial for image scaling so that model atleast sees the organ properly\n# - tissue_thickness: I am not a biology student but it is maybe how much thick the sliced tissue is. \n#                     useful: If I am doing very serious work , I might use that for splitting argumentations\n# - rle : A column which stores the mask in rle encoded format. useful: Bro, how are you going to perform masking without the mask?\n# - age : The  age of the patient\n# - sex (gender): The gender or the sex of the patient\n\n\n# Will start by exploring the orgran column\n\n\nfrom matplotlib import pyplot as plt\nfrom matplotlib import colors as mc\nimport seaborn as sns\nfrom collections import Counter\n\nsns.set_theme(\n    style=\"darkgrid\",\n    rc={\n        \"figure.facecolor\": \"#1E1E1E\",  \n        \"axes.facecolor\": \"#F5F5F5\",  \n        \"grid.color\": \"#E0E0E0\", \n        \"text.color\": \"white\", \n        \"axes.labelcolor\": \"white\", \n        \"xtick.color\": \"cyan\",  \n        \"ytick.color\": \"cyan\",  \n    },\n)\n\nclass PaletteCollections:\n\n    Rocket:  mc.ListedColormap = sns.color_palette(\"rocket\")\n    Viridis: mc.ListedColormap = sns.color_palette(\"viridis\")\n    Magma: mc.ListedColormap = sns.color_palette(\"magma\")\n    CubeHelix: mc.ListedColormap = sns.color_palette(\"cubehelix\")\n    SeaGreen: mc.ListedColormap = sns.light_palette(\"seagreen\")\n    Blues: mc.ListedColormap = sns.color_palette(\"Blues\")\n    Spectral: mc.ListedColormap = sns.color_palette(\"Spectral\")\n    Tab10: mc.ListedColormap = sns.color_palette(\"tab10\")\n    VFlag: mc.ListedColormap = sns.color_palette(\"vlag\")\n    \nfig, ax = plt.subplots(1, 1, figsize = (10, 5))\ntrain_organ_stat: dict = Counter(MetaDataSource.train_data[\"organ\"])\n\ntrain_organ_stat = dict(sorted(train_organ_stat.items(), key=lambda item: item[1]))\nsns.barplot(\n    data = None,\n    x    = train_organ_stat.keys(),\n    y    = train_organ_stat.values(),\n    hue  = train_organ_stat.keys(),\n    palette = PaletteCollections.Magma,\n    hue_order = train_organ_stat.keys(),\n    width = 0.8,\n    alpha = 0.9,\n    saturation = 0.95,\n    zorder = 3,\n    ax = ax\n)\n\nfor container in ax.containers:\n    ax.bar_label(container, padding=3)\n    \nax.grid(True)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:13.695290Z","iopub.execute_input":"2026-05-19T15:27:13.695633Z","iopub.status.idle":"2026-05-19T15:27:14.915598Z","shell.execute_reply.started":"2026-05-19T15:27:13.695610Z","shell.execute_reply":"2026-05-19T15:27:14.914435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# After getting a nice plot of the organ column , let's look at the gender distribution:\n\nfig_2, ax_2 = plt.subplots(1, 1, figsize = (10, 6))\ntrain_gender_stat: dict = Counter(MetaDataSource.train_data[\"sex\"])\n\ntrain_gender_stat = dict(sorted(train_gender_stat.items(), key=lambda item: item[1]))\nsns.barplot(\n    data = None,\n    x    = train_gender_stat.keys(),\n    y    = train_gender_stat.values(),\n    hue  = train_gender_stat.keys(),\n    palette = PaletteCollections.Tab10,\n    hue_order = train_gender_stat.keys(),\n    width = 0.8,\n    alpha = 0.9,\n    saturation = 0.95,\n    zorder = 3,\n    ax = ax_2\n)\n\n    \nax_2.grid(True)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:18.090213Z","iopub.execute_input":"2026-05-19T15:27:18.090826Z","iopub.status.idle":"2026-05-19T15:27:18.204776Z","shell.execute_reply.started":"2026-05-19T15:27:18.090791Z","shell.execute_reply":"2026-05-19T15:27:18.203331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n\n\nage_max: float = MetaDataSource.train_data['age'].max()\nage_min: float = MetaDataSource.train_data['age'].min()\n\nnum_bins: int = 5\ncounts: list[int] = [0 for _ in range(num_bins)]\n\nbin_width = (age_max - age_min) / num_bins\nbin_edges = [age_min + i * bin_width for i in range(num_bins + 1)]\n\nfor age in MetaDataSource.train_data['age']:\n\n    if age == age_max:\n        counts[-1] += 1\n        continue\n        \n    for i in range(num_bins):\n        if bin_edges[i] <= age < bin_edges[i+1]:\n            counts[i] += 1\n            break\n\nbin_labels = []\nfor i in range(num_bins):\n    bin_labels.append(f\"{int(bin_edges[i])}-{int(bin_edges[i+1])}\")\n\n\nplt.figure(figsize=(10, 6))\n\nplt.bar(bin_labels, counts, color='#3498db', edgecolor='#2c3e50', width=0.6)\n\n\nplt.title('Patient Age Distribution (HuBMAP 2022 Dataset)', fontsize=14, pad=15, fontweight='bold')\nplt.xlabel('Age Group Ranges (Years)', fontsize=12, labelpad=10)\nplt.ylabel('Patient/Sample Volume Count', fontsize=12, labelpad=10)\nplt.grid(axis='y', linestyle='--', alpha=0.5)\n\n\nfor i, count in enumerate(counts):\n    plt.text(i, count + 2, str(count), ha='center', va='bottom', fontsize=11, fontweight='bold')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:21.302403Z","iopub.execute_input":"2026-05-19T15:27:21.302721Z","iopub.status.idle":"2026-05-19T15:27:21.449244Z","shell.execute_reply.started":"2026-05-19T15:27:21.302698Z","shell.execute_reply":"2026-05-19T15:27:21.448465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Perfect , now we have a good grapse on the meta data ,let's actually do some imaging\n\n\n# Step - 1: We need to make a utility section for decoding rle into actual images\ndef decode_rle(\n    rle_str: str ,\n    mask_resolution: tuple[int, int]) -> np.ndarray:\n    \n    decoded: np.ndarray = np.zeros(shape = mask_resolution[0]*mask_resolution[1], dtype = np.uint8)\n    \n    starts_lens = list(map(int , rle_str.split(' ')))\n    \n    starts = np.asarray(starts_lens[::2], dtype = int) - 1\n    lens = np.asarray(starts_lens[1::2], dtype = int)\n    ends = starts + lens\n    \n    for start, end in zip(starts, ends):\n        decoded[start:end] = 1\n        \n    return decoded.reshape(mask_resolution, order = \"F\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:28.141845Z","iopub.execute_input":"2026-05-19T15:27:28.142186Z","iopub.status.idle":"2026-05-19T15:27:28.149138Z","shell.execute_reply.started":"2026-05-19T15:27:28.142157Z","shell.execute_reply":"2026-05-19T15:27:28.147777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Nice , we can easily visualize now:\nfrom random import randint as RandomInt\nimport cv2\n\n\nviz_shape = (2, 2)\nfig_3, axes_3 = plt.subplots(*viz_shape)\n\nfor j in range(viz_shape[1]):\n    \n    seed = RandomInt(0, MetaDataSource.train_data.shape[0])\n    row = MetaDataSource.train_data.iloc[seed]\n        \n    image_path: str = os.path.join(DataPaths.train_images_paths, f\"{row[\"id\"]}.tiff\")\n    image: np.ndarray = cv2.imread(image_path, cv2.IMREAD_COLOR_RGB)\n    mask: np.ndarray = decode_rle(row['rle'], mask_resolution = image.shape[:2])\n    axes_3[0, j].imshow(image, cmap = \"viridis\")\n    axes_3[1, j].imshow(mask, cmap = \"grey\")\n    \n    axes_3[0, j].grid(False)\n    axes_3[1, j].grid(False)\n\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:31.565480Z","iopub.execute_input":"2026-05-19T15:27:31.565868Z","iopub.status.idle":"2026-05-19T15:27:35.345131Z","shell.execute_reply.started":"2026-05-19T15:27:31.565841Z","shell.execute_reply":"2026-05-19T15:27:35.344057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Second , we will make a image preprocesser (extract images with 512 x 512 tiles , stride = 25\nfrom tqdm import tqdm as progress_bar\n\n\nclass IPConfig:\n    imgs_root_path: str = DataPaths.train_images_paths\n    meta_data: pd.DataFrame = MetaDataSource.train_data\n    target_pixel_size: float = 0.5\n    img_out_path: str = \"/kaggle/working/tiled_images\"\n    mask_out_path: str = \"/kaggle/working/tiled_masks\"\n    tile_shape: tuple[int, int] = (512, 512)\n    tile_stride: tuple[int, int] = (256, 256)\n    indices: list[int] = [] # useful for batchwise \n\n\n_DEFAULT_IP_CONFIG: IPConfig = IPConfig()\nclass ImageProcessor:\n\n    def __init__(self, config: IPConfig = _DEFAULT_IP_CONFIG):\n        os.makedirs(config.img_out_path, exist_ok = True)\n        os.makedirs(config.mask_out_path, exist_ok = True)\n        \n        self.tile_shape        = config.tile_shape\n        self.tile_stride       = config.tile_stride\n        self.imgs_root_path    = config.imgs_root_path\n        self.img_out_path      = config.img_out_path\n        self.mask_out_path     = config.mask_out_path\n        self.target_pixel_size = config.target_pixel_size\n        self.meta_data = config.meta_data.iloc[config.indices].sort_values(by = \"id\")\n        \n    def _rescale_image(self, \n                       img, \n                       curr_ps: float,\n                       image_type: str = \"image\") -> np.ndarray:\n        \n        img_h, img_w = img.shape[:2]\n        scale: float = self.target_pixel_size/curr_ps\n        img_new_shape = (int(round(img_w*scale)), int(round(img_h*scale)))\n\n        interpolation = cv2.INTER_LINEAR if image_type == \"image\" else cv2.INTER_NEAREST\n        return cv2.resize(img, img_new_shape, interpolation = cv2.INTER_LINEAR)\n        \n    def _extract_tiles_and_save(\n        self, \n        img: np.ndarray,\n        out_path: str,\n        img_type: str = \"image\"\n    ) -> int:\n        img_h, img_w = img.shape[:2]\n\n        pad_h = (self.tile_stride[0] - (img_h - self.tile_shape[0]) % self.tile_stride[0]) % self.tile_stride[0]\n        pad_w = (self.tile_stride[1] - (img_w - self.tile_shape[1]) % self.tile_stride[1]) % self.tile_stride[1]\n\n        if img_type == \"image\":\n            padded_img = np.pad(img, ((0, pad_h), (0, pad_w), (0, 0)), mode='constant', constant_values=0)\n        elif img_type == \"mask\":\n            padded_img = np.pad(img, ((0, pad_h), (0, pad_w)), mode='constant', constant_values=0)\n        else: \n            raise ValueError(f\"Unsupported img_type: {img_type}\")\n            \n        p_h, p_w = padded_img.shape[:2]\n        tile_counter: int = 0\n\n        for y in range(0, p_h - self.tile_shape[0] + 1, self.tile_stride[0]):\n            for x in range(0, p_w - self.tile_shape[1] + 1, self.tile_stride[1]):\n                \n                if img_type == \"image\":\n                    tile = padded_img[y:y + self.tile_shape[0], x:x + self.tile_shape[1], :]\n                else:\n                    tile = padded_img[y:y + self.tile_shape[0], x:x + self.tile_shape[1]]\n                \n                file_name = f\"{tile_counter}.npy\"\n                full_out_path = os.path.join(out_path, file_name)\n                \n                if img_type == \"image\":\n                    np.save(full_out_path, tile)\n                elif img_type == \"mask\":\n                    packed_tile = np.packbits(tile)\n                    np.save(full_out_path, packed_tile)\n                    \n                tile_counter += 1\n                \n        return tile_counter \n\n    def process(self) -> None:\n        pbar = progress_bar(self.meta_data.iterrows())\n\n        for index, row in pbar:\n            \n            image_path: str = os.path.join(self.imgs_root_path, f\"{row[\"id\"]}.tiff\")\n            tiled_image_folder: str = os.path.join(self.img_out_path, f\"image_{row[\"id\"]}\")\n            tiled_mask_folder: str = os.path.join(self.mask_out_path, f\"mask_{row[\"id\"]}\")\n            os.makedirs(tiled_image_folder, exist_ok = True)\n            os.makedirs(tiled_mask_folder, exist_ok = True)\n            \n            image: np.ndarray = cv2.imread(image_path, cv2.IMREAD_COLOR_RGB)\n            mask: np.ndarray = decode_rle(row[\"rle\"], mask_resolution = image.shape[:2])\n            \n            image = self._rescale_image(img = image, curr_ps = row[\"pixel_size\"], image_type = \"image\")\n            mask  = self._rescale_image(img = mask, curr_ps = row[\"pixel_size\"], image_type = \"mask\")\n\n            self._extract_tiles_and_save(image, \n                                         out_path = tiled_image_folder,\n                                         img_type = \"image\")\n            self._extract_tiles_and_save(mask, \n                                         out_path = tiled_mask_folder,\n                                         img_type = \"mask\")\n            pbar.set_postfix({\n                \"Processing\" : row[\"id\"]\n            })\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:27:43.024803Z","iopub.execute_input":"2026-05-19T15:27:43.025083Z","iopub.status.idle":"2026-05-19T15:27:43.045881Z","shell.execute_reply.started":"2026-05-19T15:27:43.025059Z","shell.execute_reply":"2026-05-19T15:27:43.044763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# we will split all the images into 3 groups , will stratify with organs so that all batches are uniform\nfrom sklearn.model_selection import train_test_split\n\ndef split_indices(config: IPConfig, batch_count: int = 3) -> dict[str, list[list]]:\n    indices_split: dict[str, tuple[list, list]] = {}\n    organs = config.meta_data[\"organ\"].tolist()\n\n    for organ in config.meta_data[\"organ\"].unique():\n        \n        organ_count: int = organs.count(organ)\n        all_indices: list[int] = [index for index, value in enumerate(organs) if value == organ]\n        \n        start: int = 0\n        num_items: int = organ_count//batch_count\n        out_list: list = []\n        \n        while start < organ_count:\n            out_list.append(all_indices[start : min(start + num_items, organ_count)])\n            start += num_items\n    \n        if (len(out_list) - batch_count) == 1:\n            for item in out_list[batch_count]:\n                out_list[batch_count - 1].append(item) # add to the previous buffer from the last last buffer\n            out_list.pop()\n            \n        indices_split[organ] = out_list\n    return indices_split\n\n\ndef merge_ow_indices(organwise_split: dict, batch_count: int):\n    out_list: list[list] = [[] for _ in range(batch_count)]\n    for _, val in organwise_split.items():\n        for idx, item in enumerate(val):\n            out_list[idx].extend(item)\n    return out_list","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:53:00.969098Z","iopub.execute_input":"2026-05-19T15:53:00.969439Z","iopub.status.idle":"2026-05-19T15:53:00.978953Z","shell.execute_reply.started":"2026-05-19T15:53:00.969412Z","shell.execute_reply":"2026-05-19T15:53:00.977753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageProcessors:\n\n    config_1 = IPConfig()\n    config_2 = IPConfig()\n    config_3 = IPConfig()\n\n\n    config_1.batch_index = 0\n    config_2.batch_index = 1\n    config_3.batch_index = 2\n\n    config_1.indices = split_indices(config_1, 3)\n    config_2.indices = split_indices(config_2, 3)\n    config_3.indices = split_indices(config_3, 3)\n\n    config_1.indices = merge_ow_indices(config_1.indices, 3)[0]\n    config_2.indices = merge_ow_indices(config_2.indices, 3)[1]\n    config_3.indices = merge_ow_indices(config_3.indices, 3)[1]\n\n    imp_1: ImageProcessor = ImageProcessor(config_1)\n    imp_2: ImageProcessor = ImageProcessor(config_1)\n    imp_3: ImageProcessor = ImageProcessor(config_1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T15:53:29.201054Z","iopub.execute_input":"2026-05-19T15:53:29.201354Z","iopub.status.idle":"2026-05-19T15:53:29.212191Z","shell.execute_reply.started":"2026-05-19T15:53:29.201328Z","shell.execute_reply":"2026-05-19T15:53:29.211247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ImageProcessors.imp_1.process()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T16:02:57.617411Z","iopub.execute_input":"2026-05-19T16:02:57.617735Z","iopub.status.idle":"2026-05-19T16:08:21.040170Z","shell.execute_reply.started":"2026-05-19T16:02:57.617711Z","shell.execute_reply":"2026-05-19T16:08:21.038791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Alright , now we will make the data loader for our \n\n\nfrom torch.utils.data import Dataset\n\nimport os\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom sklearn.preprocessing import LabelEncoder\n\nclass Datastream(Dataset):\n\n    def __init__(self, \n                 imp: ImageProcessor, \n                 indices: list[tuple[str, str, int]],\n                 transfrom = None):\n        super().__init__()\n        self.imp = imp\n        self.indices = indices\n        self.transfrom = transfrom\n\n        self.organ_encoder = LabelEncoder()\n        self.gender_encoder = LabelEncoder()\n\n        self.organs = self.organ_encoder.fit_transform(self.imp.meta_data[\"organ\"].tolist())\n        self.genders = self.gender_encoder.fit_transform(self.imp.meta_data[\"gender\"].tolist())\n        self.tissue_tickness = self.imp.meta_data[\"tissue_thickness\"].tolist()\n        self.ages = self.imp.meta_data[\"ages\"].tolist()\n        \n    def __len__(self) -> int:\n        return len(self.indices)\n\n    def __getitem__(self, index: int) -> tuple:\n\n        img_fname, mask_fname, file_idx = self.indices[index]\n        \n        img_path = os.path.join(img_fname, f\"{file_idx}.npy\")\n        mask_path = os.path.join(mask_fname, f\"{file_idx}.npy\")\n        \n        tile_path = os.path.join(self.imp.img_out_path, img_path)\n        mask_path = os.path.join(self.imp.mask_out_path, mask_path)\n        \n        tile = np.load(tile_path)\n        mask = np.load(mask_path)\n        \n        if self.transfrom:\n            argmt = self.transform(image = tile, mask = mask)\n            tile = argmt[\"image\"]\n            mask  = argmt[\"mask\"]\n        metadata = (\n            torch.tensor(self.organs[index], dtype=torch.long),           \n            torch.tensor(self.genders[index], dtype=torch.long),          \n            torch.tensor(self.tissue_tickness[index], dtype=torch.float), \n            torch.tensor(self.ages[index], dtype=torch.float)             \n        )\n        \n        return tile, mask, (self.organs[index], \n                            self.genders[index], \n                            self.tissue_tickness[index], \n                            self.ages[index])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T08:13:47.500210Z","iopub.execute_input":"2026-05-19T08:13:47.500657Z","iopub.status.idle":"2026-05-19T08:13:52.152468Z","shell.execute_reply.started":"2026-05-19T08:13:47.500625Z","shell.execute_reply":"2026-05-19T08:13:52.150992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n@dataclass\nclass TTSplitSizes:\n    train_size: float = 0.8\n    test_size: float = 0.2\n\n\ndef train_val_split(imp: ImageProcessor) -> tuple:\n    train_split = []\n    test_split = []\n    \n    for index, row in imp.meta_data.iterrows():\n        \n        img_path: str = os.path.join(imp.img_out_path, f\"image_{row['id']}\")\n        \n        num_tiles = len(os.listdir(img_path))\n        all_indices = [x for x in range(num_tiles)]\n        \n        train_indices, val_indices, _, _ = train_test_split(\n            all_indices,\n            all_indices, # dummy only\n            train_size = TTSplitSizes.train_size,\n            test_size = TTSplitSizes.test_size,\n            shuffle = True,\n            random_sate = RandomStates.random_state\n        )\n        \n        \n        for item in train_indices:\n            train_split.append((f\"image_{row['id']}\", f\"mask_{row['id']}\", item))\n        for item in test_indices:\n            test_split.append((f\"image_{row['id']}\", f\"mask_{row['id']}\", item))\n\n\nTrainIndices = train_val_split(imp = ImageProcessors.imp_1)\nValIndices = train_val_split(imp = ImageProcessors.imp_1)\n\nprint(f\"{c.Fore.Cyan} -> Train indices peak: {TrainIndices[:5]}, size  = {len(TrainIndices)}\")\nprint(f\"{c.Fore.Cyan} -> Val indices peak: {ValIndices[:5]}, size = {len(TrainIndices)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}