{"metadata":{"kernelspec":{"display_name":"Flagellar_Location","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.16"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"83029388de6e597b","cell_type":"markdown","source":"# Flagellar Location\nIn this project we want to predict the location of flagellar proteins in tomographies of bacteria.\n\nTo achieve this, we will first extract from the tomographies different blocks, and with those we will train a convolutional neural network (CNN) to predict if the flagellar protein is present in a specific block or not.\n\nThen we will use the trained CNN to predict the presence of a flagellar protein in different blocks of the tomographies. And if it's present in one of them, we will iterate by subdividing the block in other smaller blocks, and repeat the process.\n\n","metadata":{}},{"id":"4d1d875f","cell_type":"markdown","source":"## 1 Notebook Setup\nIn this first part we will set up all the necessary libraries and functions to run this project.","metadata":{}},{"id":"185e25281e7df5a3","cell_type":"markdown","source":"### 1.1 Library Imports and Environment Setup\nFirst we need to import the libraries we are going to use in this project.","metadata":{}},{"id":"e5613279263e2274","cell_type":"code","source":"import plotly.graph_objects as go\nimport plotly.io as pio\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport os\nimport random\nimport torch\nimport torch.nn.functional as F\nimport warnings\nimport threading\nfrom itertools import cycle\nimport imageio.v2 as imageio\nimport numpy as np\nimport pandas as pd\nimport os\nimport random\nimport torch\nimport imageio.v2 as imageio\nfrom tqdm import tqdm\nimport sys\nimport threading\nimport traceback\nimport pickle\nimport hashlib\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score\nimport torch.nn.functional as F\nimport random\nfrom torch.utils.data import WeightedRandomSampler\nfrom mpl_toolkits.mplot3d import Axes3D\nimport plotly.graph_objects as go\nimport plotly.io as pio\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nfrom mpl_toolkits.mplot3d import Axes3D \nimport concurrent.futures\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom typing import List, Tuple\n\n\nos.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'","metadata":{"ExecuteTime":{"end_time":"2025-05-26T12:41:27.955811Z","start_time":"2025-05-26T12:41:24.570729Z"}},"outputs":[],"execution_count":null},{"id":"cde3bea0","cell_type":"markdown","source":"### 1.2 Tomography utils\nWe define here 3 utils classes for better organization and type checking, Point3D, Box3D, BoxShape","metadata":{}},{"id":"3f0d84a1","cell_type":"code","source":"class Point3D:\n    \"\"\"Represents a 3D point with x, y, z coordinates\"\"\"\n    \n    def __init__(self, x, y, z):\n        self.x = x\n        self.y = y\n        self.z = z\n    \n    def to_zyx_tuple(self):\n        \"\"\"Convert to (z, y, x) tuple format used in existing code\"\"\"\n        return (self.z, self.y, self.x)\n    \n    def __str__(self):\n        return f\"Point3D(x: {self.x}, y: {self.y}, z: {self.z})\"\n    \n    def __repr__(self):\n        return f\"Point3D({self.x}, {self.y}, {self.z})\"\n    \n    def __eq__(self, other):\n        if not isinstance(other, Point3D):\n            return False\n        return self.x == other.x and self.y == other.y and self.z == other.z\n    \n\nclass Box3D:\n    \n    def __init__(self, x_start, x_end, y_start, y_end, z_start, z_end):\n        self.x_start = x_start\n        self.x_end = x_end\n        self.y_start = y_start\n        self.y_end = y_end\n        self.z_start = z_start\n        self.z_end = z_end\n    \n    def to_ranges(self):\n        \"\"\"Convert to tuple ranges format: ((x_start, x_end), (y_start, y_end), (z_start, z_end))\"\"\"\n        return ((self.x_start, self.x_end), (self.y_start, self.y_end), (self.z_start, self.z_end))\n    \n    def get_dimensions(self):\n        \"\"\"Get cube dimensions (width, height, depth)\"\"\"\n        width = self.x_end - self.x_start\n        height = self.y_end - self.y_start\n        depth = self.z_end - self.z_start\n        return (width, height, depth)\n    \n    def get_center(self):\n        \"\"\"Get center point of the cube as Point3D\"\"\"\n        center_x = (self.x_start + self.x_end) / 2\n        center_y = (self.y_start + self.y_end) / 2\n        center_z = (self.z_start + self.z_end) / 2\n        return Point3D(center_x, center_y, center_z)\n    \n    def contains_point(self, point : Point3D):\n        \"\"\"Check if a Point3D is inside this cube\"\"\"\n        if not isinstance(point, Point3D):  \n            raise TypeError(\"Expected a Point3D instance\")\n        \n        return (self.x_start <= point.x <= self.x_end and\n                self.y_start <= point.y <= self.y_end and\n                self.z_start <= point.z <= self.z_end)\n    \n    def __str__(self):\n        return f\"CubeCoords(x: {self.x_start}-{self.x_end}, y: {self.y_start}-{self.y_end}, z: {self.z_start}-{self.z_end})\"\n    \n    def __repr__(self):\n        return f\"CubeCoords({self.x_start}, {self.x_end}, {self.y_start}, {self.y_end}, {self.z_start}, {self.z_end})\"\n    \nclass BoxShape:\n    \n    def __init__(self, center, x, y, z):\n        self.center = center\n        self.x = x\n        self.y = y\n        self.z = z\n    \n    def get_box_coords(self):\n        \"\"\"Get Box2D coordinates based on center and dimensions\"\"\"\n        x_start = self.center.x - self.x / 2\n        x_end = self.center.x + self.x / 2\n        y_start = self.center.y - self.y / 2\n        y_end = self.center.y + self.y / 2\n        z_start = self.center.z - self.z / 2\n        z_end = self.center.z + self.z / 2\n        return Box3D(x_start, x_end, y_start, y_end, z_start, z_end)\n    \n    def toTuple(self):\n        return (self.x, self.y, self.z)","metadata":{},"outputs":[],"execution_count":null},{"id":"c77bdeba2b9227a0","cell_type":"markdown","source":"### 1.3 Hyperparameters\nHere the list for all the hyperparameters used for this notebook","metadata":{}},{"id":"1aa13f337e4bd60a","cell_type":"code","source":"EXPORT_SHAPE = BoxShape(x=30, y=30, z=30, center=None)\nEXTRACTION_ITERATIONS = 2\nANALYSIS_ITERATIONS = 7\nEXTRACTION_PADDING = 2\nEXPORT_WORKERS = 1\nDEBUG = False\nANTIALIASING = False\n\nMINSIZE_RATIO = 0.4\nFIRST_MINSIZE_RATIO = 1\n\nDROPOUT_RATE = 0.5\nBATCH_SIZE = 128\nEPOCHS = 128\nLR = 0.002\nGAMMA = 0.99\nSEED = 42\nVALIDATION_SET_SIZE = 0.25\nTRAINING_NOISE = 0.0005\nTRAINING_MIRROR_PROB = 0.5\nPROBABILITY_TRESHOLD = 0.5\nALPHA_RATIO = 1","metadata":{"ExecuteTime":{"end_time":"2025-05-26T12:41:27.997911Z","start_time":"2025-05-26T12:41:27.987127Z"}},"outputs":[],"execution_count":null},{"id":"85664b76","cell_type":"markdown","source":"### 1.4 Seeding and Device Setup","metadata":{}},{"id":"4df4367c","cell_type":"code","source":"# --- Seeding ---\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\n# --- Detecting if on Kaggle ---\nON_KAGGLE = True if 'KAGGLE_URL_BASE' in os.environ else False\nprint(f\"Running on Kaggle: {ON_KAGGLE}\")\n\n# --- Device Setup ---\ndevice_ids = []\nnum_gpus = 0\nDEFAULT_DEVICE = \"cpu\"\nDEVICE = DEFAULT_DEVICE\n\nif torch.cuda.is_available():\n    num_gpus = torch.cuda.device_count()\n    if num_gpus > 0:\n        print(f\"CUDA available: {torch.cuda.is_available()}\")\n        print(f\"Number of GPUs: {num_gpus}\")\n        device_ids = [f'cuda:{i}' for i in range(num_gpus)] \n        for i in range(num_gpus):\n            print(f\"GPU {i}: {torch.cuda.get_device_name(i)}\")\n        DEVICE = device_ids[0] \n        torch.cuda.manual_seed_all(SEED) \n        torch.backends.cudnn.deterministic = True\n    else:\n        print(\"CUDA driver found, but no compatible GPUs detected.\")\nelse:\n    print(\"CUDA not available, using CPU.\")\n\n\n# --- GPU Assignment for Multi-GPU Export ---\ngpu_semaphores = {}\nif len(device_ids) > 1:\n    for device_id in device_ids:\n        gpu_semaphores[device_id] = threading.Semaphore(1)\n    print(f\"Created GPU semaphores for exclusive access: {list(gpu_semaphores.keys())}\")\n    \n    EXPORT_WORKERS = len(device_ids)\n    print(f\"Adjusted EXPORT_WORKERS to {EXPORT_WORKERS} to match GPU count\")\nelse:\n    print(f\"Running on single device ({DEVICE}). No GPU semaphores needed.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"138a53da73b3bba","cell_type":"markdown","source":"### 1.5 Dataset Loading","metadata":{}},{"id":"2958bc76bbe4ab6a","cell_type":"code","source":"base_dir='../input/byu-locating-bacterial-flagellar-motors-2025'\ntrain_csv_path = os.path.join(base_dir, 'train_labels.csv')\ntrain_df = pd.read_csv(train_csv_path)","metadata":{"ExecuteTime":{"end_time":"2025-05-26T12:41:28.080113Z","start_time":"2025-05-26T12:41:28.068273Z"}},"outputs":[],"execution_count":null},{"id":"f5ba35b8dff9c7b4","cell_type":"markdown","source":"### 1.6 Tomography Dataset Export and Transformation Functions\nIn this next big cell I'm going to reimplement the resize function with **Antialiasing** optimized for cuda. I did this because I was not able to find this in any already existing library. As it's lengthy and boring I suggest to check on it at the end.\n\nNaively since showing at screen the tomography with antialiasing was looking better I thought the model would benefit training on resized block with antialiasing, instead this is not happening. The model **train better when the Antialiasing is turned OFF**. So this optimized function is used when displaying the tomography with antialiasing enabled. ","metadata":{}},{"id":"33111a6410fed0e2","cell_type":"code","source":"def calculate_sigma(factors, anti_aliasing_sigma):\n    if anti_aliasing_sigma is None:\n        threshold = 1.0 + 1e-8\n        sigma = torch.maximum(\n            torch.tensor(0.0, device=factors.device, dtype=factors.dtype),\n             (factors - 1) / 2.0\n        )\n        sigma = torch.where(factors > threshold, sigma, torch.tensor(0.0, device=sigma.device, dtype=sigma.dtype))\n    else:\n        sigma = torch.atleast_1d(torch.as_tensor(anti_aliasing_sigma, dtype=factors.dtype, device=factors.device))\n        if sigma.numel() == 1:\n             sigma = sigma.repeat(len(factors))\n        elif sigma.numel() != len(factors):\n             raise ValueError(\"anti_aliasing_sigma must have one element or match input dimensions\")\n\n        if torch.any(sigma < 0):\n            raise ValueError(\"anti_aliasing_sigma must be non-negative\")\n    return sigma\n\ndef create_gaussian_kernel_1d(sigma, kernel_size, dtype=torch.float32, device='cpu'):\n    \"\"\"Creates a 1D Gaussian kernel.\"\"\"\n    if sigma <= 1e-8:\n        if kernel_size % 2 == 0: kernel_size +=1\n        kernel = torch.zeros(kernel_size, dtype=dtype, device=device)\n        kernel[kernel_size // 2] = 1.0\n        return kernel\n\n    coords = torch.arange(kernel_size, dtype=dtype, device=device)\n    center = kernel_size // 2\n    sigma_sq = sigma**2\n    kernel = torch.exp(-(coords - center)**2 / (2 * sigma_sq + 1e-12))\n    kernel = kernel / kernel.sum()\n    return kernel\n\ndef gaussian_filter_3d(image_tensor, sigma_tensor, mode='reflect', cval=0.0):\n    \"\"\"Applies 3D Gaussian filtering using separable 1D convolutions.\"\"\"\n    dtype = image_tensor.dtype\n    device = image_tensor.device\n    spatial_dims = image_tensor.ndim - 2\n    channels = image_tensor.shape[1]\n\n    if not torch.any(sigma_tensor > 1e-4):\n        return image_tensor\n\n    truncate = 4.0\n\n    kernel_sizes = (torch.ceil(sigma_tensor.to(dtype=torch.float32) * truncate) * 2 + 1).int()\n\n    filtered_tensor = image_tensor\n    for i in range(spatial_dims):\n        s = sigma_tensor[i]\n        kernel_size = kernel_sizes[i].item() \n\n        if s <= 1e-4 or kernel_size <= 1:\n            continue\n\n        kernel_1d = create_gaussian_kernel_1d(s.item(), kernel_size, dtype=dtype, device=device)\n        kernel_reshaped = kernel_1d.reshape(1, 1, -1)\n\n        conv_padding = (kernel_size - 1) // 2\n\n        if i == 0: \n            kernel_conv = kernel_reshaped.reshape(1, 1, int(kernel_size), 1, 1).repeat(channels, 1, 1, 1, 1)\n            conv_func = F.conv3d\n            padding_arg = (int(conv_padding), 0, 0)\n        elif i == 1: \n            kernel_conv = kernel_reshaped.reshape(1, 1, 1, int(kernel_size), 1).repeat(channels, 1, 1, 1, 1)\n            conv_func = F.conv3d\n            padding_arg = (0, int(conv_padding), 0)\n        else: \n            kernel_conv = kernel_reshaped.reshape(1, 1, 1, 1, int(kernel_size)).repeat(channels, 1, 1, 1, 1)\n            conv_func = F.conv3d\n            padding_arg = (0, 0, int(conv_padding))\n\n        pad_arg_tuple = (int(padding_arg[2]), int(padding_arg[2]) if i==2 else 0,\n                         int(padding_arg[1]), int(padding_arg[1]) if i==1 else 0,\n                         int(padding_arg[0]), int(padding_arg[0]) if i==0 else 0)\n\n        if mode in ['reflect', 'constant']: \n             pad_mode = 'reflect' if mode == 'reflect' else 'constant'\n             pad_value = cval if mode == 'constant' else 0.0\n             padded_tensor = F.pad(filtered_tensor, pad_arg_tuple, mode=pad_mode, value=pad_value)\n             conv_padding_applied = 0\n             filtered_tensor = conv_func(padded_tensor, kernel_conv, padding=conv_padding_applied, groups=channels)\n\n        else: \n            filtered_tensor = conv_func(filtered_tensor, kernel_conv, padding=padding_arg, groups=channels)\n\n    return filtered_tensor\n\n\ndef interpolate(image_tensor, output_shape_spatial, order, align_corners):\n    \"\"\" Maps spline order to PyTorch interpolation mode and performs interpolation \"\"\"\n    if order == 0:\n        interp_mode = 'nearest'\n    elif order == 1:\n        interp_mode = 'trilinear'\n    elif order >= 2 and order <= 5:\n        interp_mode = 'trilinear' \n    else:\n        interp_mode = 'trilinear'\n\n    if image_tensor.shape[2:] == tuple(output_shape_spatial):\n        return image_tensor\n\n    return F.interpolate(image_tensor,\n                         size=tuple(output_shape_spatial), \n                         mode=interp_mode,\n                         align_corners=align_corners\n                         )\n\n\ndef resize_cuda_optimized(\n    image: np.ndarray,\n    output_shape: tuple,\n    order: int | None = None,\n    mode: str = 'reflect',\n    cval: float = 0.0,\n    clip: bool = True, \n    preserve_range: bool = False, \n    anti_aliasing: bool | None = None,\n    anti_aliasing_sigma: float | tuple[float] | None = None,\n    device: str | None = None\n):\n    \"\"\"\n    Resize N-dimensional image using PyTorch, mimicking skimage.transform.resize.\n    Optimized for memory efficiency and direct uint8 output when preserve_range=True.\n    \"\"\"\n    # --- Start: Setup code (mostly unchanged) ---\n    input_shape = image.shape\n    input_dtype = image.dtype\n    input_ndim = image.ndim\n\n    output_shape = tuple(output_shape)\n    output_ndim = len(output_shape)\n\n    if output_ndim > input_ndim:\n        input_shape_adjusted = input_shape + (1,) * (output_ndim - input_ndim)\n        image_adjusted = np.reshape(image, input_shape_adjusted)\n        input_ndim_adjusted = image_adjusted.ndim\n    elif output_ndim == input_ndim - 1 and input_ndim > 1:\n         if input_shape[-1] > 4:\n             warnings.warn(f\"Input shape {input_shape} last dimension is large, potentially not channels. Treating as spatial.\")\n             raise ValueError(f\"Output shape {output_shape} has fewer dimensions than input {input_shape}, and last input dim is large.\")\n         else:\n             output_shape = output_shape + (input_shape[-1],)\n             image_adjusted = image\n             input_shape_adjusted = input_shape\n             input_ndim_adjusted = input_ndim\n             output_ndim += 1\n    elif output_ndim == input_ndim:\n        image_adjusted = image\n        input_shape_adjusted = input_shape\n        input_ndim_adjusted = input_ndim\n    else:\n        raise ValueError(f\"output_shape length {output_ndim} must be >= image.ndim {input_ndim} or == image.ndim - 1\")\n\n    if order is None: order = 0 if input_dtype == bool else 1\n    if not (0 <= order <= 5): raise ValueError(\"Order must be in range 0-5.\")\n    if input_dtype == bool and order != 0: raise ValueError(\"Cannot use order > 0 for bool images.\")\n\n    is_downsampling = any(x > y for x, y in zip(input_shape_adjusted, output_shape))\n    if anti_aliasing is None:\n        anti_aliasing = (\n            not input_dtype == bool\n            and not (np.issubdtype(input_dtype, np.integer) and order == 0)\n            and is_downsampling\n        )\n    if input_dtype == bool and anti_aliasing:\n        raise ValueError(\"anti_aliasing must be False for boolean images\")\n\n    input_range = None\n    if clip and not preserve_range:\n        if input_dtype == bool: input_range = (False, True)\n        elif np.issubdtype(input_dtype, np.integer):\n            try: iinfo = np.iinfo(input_dtype); input_range = (iinfo.min, iinfo.max)\n            except ValueError: pass\n        elif np.issubdtype(input_dtype, np.floating):\n             input_range = (np.min(image), np.max(image))\n\n    process_dtype = torch.float32\n    if image_adjusted.dtype == np.float16:\n        image_adjusted = image_adjusted.astype(np.float32)\n\n\n    with torch.no_grad():\n        try:\n            image_tensor = torch.from_numpy(image_adjusted).to(device)\n        except Exception as e:\n            print(f\"Error converting/moving NumPy array (shape {image_adjusted.shape}, dtype {image_adjusted.dtype}) to {device}: {e}\")\n            raise\n\n        image_tensor = image_tensor.to(process_dtype)\n\n        if not preserve_range and (input_dtype == bool or np.issubdtype(input_dtype, np.integer)):\n             img_min = float(image_tensor.min())\n             img_max = float(image_tensor.max())\n             if img_max > img_min:\n                  image_tensor = ((image_tensor.double() - img_min) / (img_max - img_min)).to(process_dtype)\n\n        channel_dim_index = -1\n        spatial_dims = input_ndim_adjusted\n        if input_ndim_adjusted > 1 and input_shape_adjusted[-1] == output_shape[-1] and input_shape_adjusted[-1] <= 4: \n             is_multichannel = (len(output_shape) == input_ndim_adjusted) or \\\n                               (len(output_shape) -1 == input_ndim_adjusted -1 and output_ndim == input_ndim_adjusted)\n             if is_multichannel:\n                  channel_dim_index = input_ndim_adjusted - 1\n                  spatial_dims = input_ndim_adjusted - 1\n\n        if channel_dim_index == -1: \n             image_tensor = image_tensor.unsqueeze(0).unsqueeze(0) \n             output_shape_spatial = output_shape\n        else: \n             permute_dims = [channel_dim_index] + list(range(channel_dim_index))\n             image_tensor = image_tensor.permute(*permute_dims).unsqueeze(0) \n             output_shape_spatial = output_shape[:-1]\n\n        output_shape_spatial = tuple(int(x) for x in output_shape_spatial)\n\n        filtered_tensor = image_tensor\n        if anti_aliasing:\n            current_spatial_shape = image_tensor.shape[2:]\n            factors_np = np.array([\n                i / o if o > 0 else 1.0 \n                for i, o in zip(current_spatial_shape, output_shape_spatial)\n            ])\n            factors = torch.as_tensor(factors_np, dtype=process_dtype, device=device)\n            sigma = calculate_sigma(factors, anti_aliasing_sigma)\n\n            if spatial_dims == 3:\n                filtered_tensor = gaussian_filter_3d(image_tensor, sigma, mode=mode, cval=cval)\n            else:\n                 if torch.any(sigma > 1e-4):\n                     warnings.warn(f\"Anti-aliasing filter not implemented for {spatial_dims} spatial dimensions, skipping.\")\n\n        interpolated_tensor = interpolate(filtered_tensor, output_shape_spatial, order, align_corners=False)\n        \n\n        output_tensor = interpolated_tensor.squeeze(0)\n\n        if channel_dim_index == -1: \n             if output_tensor.shape[0] == 1:\n                 output_tensor = output_tensor.squeeze(0) \n        else: \n             permute_back_dims = list(range(1, spatial_dims + 1)) + [0]\n             output_tensor = output_tensor.permute(*permute_back_dims) \n\n\n        if preserve_range:\n            output_tensor = output_tensor.clamp_(0, 255).round_().to(torch.uint8)\n            output_final = output_tensor.cpu().numpy()\n\n        else: \n             if clip and input_range is not None: \n                 min_val, max_val = input_range\n                 min_val_t = torch.tensor(min_val, dtype=output_tensor.dtype, device=device)\n                 max_val_t = torch.tensor(max_val, dtype=output_tensor.dtype, device=device)\n                 output_tensor = torch.clamp(output_tensor, min=min_val_t, max=max_val_t)\n\n             output_np = output_tensor.cpu().numpy()\n             output_final = output_np if output_np.dtype == np.float32 else output_np.astype(np.float32)\n\n\n    if output_final.shape != tuple(output_shape):\n         warnings.warn(f\"Final shape {output_final.shape} doesn't match target {tuple(output_shape)}. This might be due to interpolation precision or channel handling.\", UserWarning)\n\n    return output_final\n\ntry:\n    if hasattr(torch, 'compile'):\n        print(\"Attempting to compile resize_pytorch_v2_optimized with torch.compile(mode=\\\"reduce-overhead\\\")...\")\n        compiled_resize_pytorch_v2_optimized = torch.compile(resize_cuda_optimized, mode=\"reduce-overhead\")\n        print(\"resize_pytorch_v2_optimized compiled successfully.\")\n    else:\n        print(\"torch.compile not available (requires PyTorch 2.0+). Proceeding without compilation.\")\nexcept Exception as e:\n    print(f\"Error during torch.compile: {e}. Proceeding without compilation.\")","metadata":{"ExecuteTime":{"end_time":"2025-05-26T12:41:30.275166Z","start_time":"2025-05-26T12:41:28.212851Z"}},"outputs":[],"execution_count":null},{"id":"175727b8cfb0ba3","cell_type":"markdown","source":"### 1.7 Manipulating Tomography Functions\nIn this section we define functions necessary to manipulate the tomographies while loading the train and test datasets.","metadata":{}},{"id":"3b01462c0780dca3","cell_type":"code","source":"def getImagesPath(base_dir, folder, tomo_id):\n\n    folder_path = os.path.join(base_dir, folder, tomo_id)\n    image_files = sorted([\n        os.path.join(folder_path, f)\n        for f in os.listdir(folder_path)\n        if f.endswith('.jpg')\n    ])\n    return image_files\n\ndef load_tomography(image_files, dtype=np.uint16) -> np.ndarray:\n    first_img = imageio.imread(image_files[0])\n    shape = (len(image_files),) + first_img.shape\n\n    tomography = np.empty(shape, dtype=dtype)\n    tomography[0] = first_img\n\n    for i, file in enumerate(image_files[1:], start=1):\n        tomography[i] = imageio.imread(file)\n\n    return tomography\n\ndef getTomoRows(df, identifier):\n    return df[df['tomo_id'] == identifier]\n\ndef getTomographyShape(train_df, tomo_id) -> BoxShape:\n    tomo_row = getTomoRows(train_df, tomo_id).iloc[0]\n    z = int(tomo_row['Array shape (axis 0)'])\n    y = int(tomo_row['Array shape (axis 1)'])\n    x = int(tomo_row['Array shape (axis 2)'])\n    \n    center = Point3D(x=x/2, y=y/2, z=z/2)\n    \n    return BoxShape(x=x, y=y, z=z, center=center)\n\ndef getFlagellarCoordinates(train_df, tomo_id) -> List[Point3D]:\n    flagellar_points = []\n    tomo_rows = getTomoRows(train_df, tomo_id)\n    for _, row in tomo_rows.iterrows():\n        point = Point3D(\n            x=int(row['Motor axis 2']),\n            y=int(row['Motor axis 1']),\n            z=int(row['Motor axis 0'])\n        )\n        flagellar_points.append(point)\n    return flagellar_points\n\ndef getCubeCoordsToExtract(original_tomography_shape : BoxShape, dots : list[Point3D], minSizeRatio:float = FIRST_MINSIZE_RATIO, padding=EXTRACTION_PADDING, export_zone: Box3D | None = None) -> List[Tuple[Box3D, bool, List[Point3D], Box3D]]:\n    \n    if original_tomography_shape is None or not isinstance(original_tomography_shape, BoxShape):\n        raise ValueError(\"original_tomography_shape must be a BoxShape object\")\n    if not isinstance(dots, list) or not all(isinstance(dot, Point3D) for dot in dots):\n        raise ValueError(\"dots must be a list of Point3D objects\")\n    if not isinstance(export_zone, (Box3D, type(None))):\n        raise ValueError(\"export_zone must be a Box2D object or None\")\n    \n\n    if export_zone is None:\n        export_zone = Box3D(\n            x_start=0, x_end=original_tomography_shape.x,\n            y_start=0, y_end=original_tomography_shape.y, \n            z_start=0, z_end=original_tomography_shape.z\n        )\n        export_zone_offset = Point3D(0, 0, 0)\n    else:\n        if not isinstance(export_zone, Box3D):\n            raise ValueError(\"export_zone must be a CubeCoords object or a tuple in the format ((z_start, z_end), (y_start, y_end), (x_start, x_end))\")\n        \n        export_zone_offset = Point3D(export_zone.x_start, export_zone.y_start, export_zone.z_start)\n    \n    export_zone_dims = export_zone.get_dimensions()\n    export_zone_shape = (export_zone_dims[2], export_zone_dims[1], export_zone_dims[0])\n    \n    original_tomography_shape_tuple = (original_tomography_shape.z, original_tomography_shape.y, original_tomography_shape.x)\n    \n    for i, (export_dim, original_dim) in enumerate(zip(export_zone_shape, original_tomography_shape_tuple)):\n        if export_dim > original_dim:\n            raise ValueError(f\"export_zone_shape[{i}] ({export_dim}) cannot be larger than original_tomography_shape[{i}] ({original_dim})\")\n  \n    min_size = min(export_zone_shape)\n    cube_coords = []\n    current_iteration = 1\n\n    side_size = min_size * minSizeRatio\n\n    z_cubes = int(export_zone_shape[0]//side_size + 1)\n    y_cubes = int(export_zone_shape[1]//side_size + 1)\n    x_cubes = int(export_zone_shape[2]//side_size + 1)\n\n    z_remainder = export_zone_shape[0] % side_size\n    y_remainder = export_zone_shape[1] % side_size\n    x_remainder = export_zone_shape[2] % side_size\n\n    z_overlapping = (side_size - z_remainder) / (z_cubes - 1) if z_cubes > 1 else 0\n    y_overlapping = (side_size - y_remainder) / (y_cubes - 1) if y_cubes > 1 else 0\n    x_overlapping = (side_size - x_remainder) / (x_cubes - 1) if x_cubes > 1 else 0\n\n    for z in range(z_cubes):\n        for y in range(y_cubes):\n            for x in range(x_cubes):\n                z_start = int(z * side_size - z_overlapping * z)\n                y_start = int(y * side_size - y_overlapping * y)\n                x_start = int(x * side_size - x_overlapping * x)\n                z_end = int(z_start + side_size)\n                y_end = int(y_start + side_size)\n                x_end = int(x_start + side_size)\n                \n                cube_coords_obj = Box3D(\n                    x_start=x_start, x_end=x_end,\n                    y_start=y_start, y_end=y_end,\n                    z_start=z_start, z_end=z_end\n                )\n                \n                original_cube_coords = Box3D(\n                    x_start=x_start + export_zone_offset.x, x_end=x_end + export_zone_offset.x,\n                    y_start=y_start + export_zone_offset.y, y_end=y_end + export_zone_offset.y,\n                    z_start=z_start + export_zone_offset.z, z_end=z_end + export_zone_offset.z\n                )\n                \n                cube_dots = []\n                hasFlagellar = False\n                \n                for dot in dots:\n\n                    if original_cube_coords.contains_point(dot):\n                        hasFlagellar = True\n                        relative_dot = Point3D(\n                            dot.x - (x_start + export_zone_offset.x),\n                            dot.y - (y_start + export_zone_offset.y),\n                            dot.z - (z_start + export_zone_offset.z)\n                        )\n                        cube_dots.append(relative_dot)\n                \n                cube_coords.append((\n                    cube_coords_obj,  \n                    hasFlagellar, \n                    cube_dots,  \n                    original_cube_coords  \n                ))\n\n    side_size = side_size / 2 + padding\n    current_iteration += 1\n\n    return cube_coords\n    \ndef getTomographyExportsByTomoId(tomo_id, base_dir, device, cube_coords: List[Tuple[Box3D, bool, List[Point3D], Box3D]], output_block_shape: BoxShape = EXPORT_SHAPE, iterations: int =1, antialiasing=ANTIALIASING) -> List[Tuple[np.ndarray, bool, List[Point3D], Box3D, Box3D]]:\n\n    if not isinstance(output_block_shape, BoxShape):\n        raise ValueError(\"output_block_shape must be a BoxShape object\")\n    if not isinstance(cube_coords, list) or not all(isinstance(coord, tuple) and len(coord) == 4 for coord in cube_coords):\n        raise ValueError(\"cube_coords must be a list of tuples in the format (cube_coord, hasFlagellar, cube_dots, original_coords)\")   \n    \n    imageFilesPath = getImagesPath(base_dir, \"train\", tomo_id)\n\n    tomography = load_tomography(imageFilesPath)\n\n    return getTomographyExports(tomography, device, cube_coords=cube_coords, output_block_shape=output_block_shape,  antialiasing=antialiasing, iterations=iterations)\n\ndef getTomographyExports(tomography, device, \n                         cube_coords: List[Tuple[Box3D, bool, List[Point3D], Box3D]],\n                         output_block_shape=BoxShape(x=20, y=20, z=20, center=None),\n                         antialiasing=False,\n                         iterations=1) -> List[Tuple[np.ndarray, bool, List[Point3D], Box3D, Box3D]]:\n\n    if not isinstance(tomography, np.ndarray):\n        raise ValueError(\"tomography must be a NumPy ndarray\")\n    if not isinstance(output_block_shape, BoxShape):\n        raise ValueError(\"output_block_shape must be a BoxShape object\")\n    if not isinstance(cube_coords, list) or not all(isinstance(coord, tuple) and len(coord) == 4 for coord in cube_coords):\n        raise ValueError(\"cube_coords must be a list of tuples in the format (cube_coord, hasFlagellar, cube_dots, original_coords)\")\n    if not isinstance(iterations, int) or iterations < 1:\n        raise ValueError(\"iterations must be a positive integer\")\n\n    tomo_shape = tomography.shape\n    tomography_shape = BoxShape(\n        x=tomo_shape[2], y=tomo_shape[1], z=tomo_shape[0],\n        center=None  \n    )\n    \n    all_exports = []\n    current_cube_coords = cube_coords\n    \n    for iteration in range(iterations):\n        iteration_exports = []\n        \n        for cube_coord, hasFlagellar, cube_dots, original_coords in current_cube_coords:\n            if not isinstance(cube_coord, Box3D):\n                raise ValueError(\"First element of cube_coords must be a Box2D object\")\n            if not isinstance(hasFlagellar, bool):\n                raise ValueError(\"Second element of cube_coords must be a boolean indicating if the cube contains flagellar dots\")\n            if not isinstance(cube_dots, list) or not all(isinstance(dot, Point3D) for dot in cube_dots):\n                raise ValueError(\"Third element of cube_coords must be a list of Point3D objects representing dots in the cube\")\n            if not isinstance(original_coords, Box3D):\n                raise ValueError(\"Fourth element of cube_coords must be a Box2D object representing the original coordinates of the cube\")\n            \n            x_start, x_end = cube_coord.x_start, cube_coord.x_end\n            y_start, y_end = cube_coord.y_start, cube_coord.y_end\n            z_start, z_end = cube_coord.z_start, cube_coord.z_end\n\n            extracted_block = tomography[z_start:z_end, y_start:y_end, x_start:x_end]\n\n            original_cube_shape = extracted_block.shape\n            \n            resized_block = resize_cuda_optimized(extracted_block, output_block_shape.toTuple(), \n                                                anti_aliasing=antialiasing, preserve_range=True, device=device)\n\n            resized_block = resized_block.astype(np.uint8)\n\n            resized_dots = []\n            if cube_dots: \n                for dot in cube_dots:\n                    scale_z = float(output_block_shape.z) / original_cube_shape[0]\n                    scale_y = float(output_block_shape.y) / original_cube_shape[1]\n                    scale_x = float(output_block_shape.x) / original_cube_shape[2]\n\n                    rescaledDot = Point3D(x=dot.x * scale_x, y=dot.y * scale_y, z=dot.z * scale_z)\n                    resized_dots.append(rescaledDot)\n\n            iteration_exports.append((resized_block, hasFlagellar, resized_dots, cube_coord, original_coords))\n        \n        all_exports.extend(iteration_exports)\n        \n        if iteration < iterations - 1:\n            next_cube_coords = []\n            \n            for cube_coord, hasFlagellar, cube_dots, original_coords in current_cube_coords:\n                if hasFlagellar and cube_dots: \n                    \n                    global_flagellar_points = []\n                    for dot in cube_dots:\n                        global_dot = Point3D(\n                            x=dot.x + original_coords.x_start,\n                            y=dot.y + original_coords.y_start,\n                            z=dot.z + original_coords.z_start\n                        )\n                        global_flagellar_points.append(global_dot)\n                    \n                    subdivided_cubes = getCubeCoordsToExtract(\n                        original_tomography_shape=tomography_shape,\n                        dots=global_flagellar_points,\n                        minSizeRatio=MINSIZE_RATIO,\n                        padding=EXTRACTION_PADDING,\n                        export_zone=original_coords\n                    )\n                    \n                    next_cube_coords.extend(subdivided_cubes)\n            \n            current_cube_coords = next_cube_coords\n            \n            if not current_cube_coords:\n                break\n\n    return all_exports","metadata":{"ExecuteTime":{"end_time":"2025-05-26T12:41:30.402008Z","start_time":"2025-05-26T12:41:30.38235Z"}},"outputs":[],"execution_count":null},{"id":"57d83167309bd34c","cell_type":"markdown","source":"### 1.8 Function Validation and Visualization Demo\nIn this block we will create some functions helping us to visualize the tomographies. Also we will use them to better understand the data we are working with.","metadata":{}},{"id":"1d5ca8ecb5dfb150","cell_type":"code","source":"def display_tomography(tomography: np.ndarray, title, flagellar_points: List[Point3D] | None = None):\n\n    if not isinstance(tomography, np.ndarray):\n        raise ValueError(\"tomography must be a numpy array\")\n    if not isinstance(flagellar_points, (list, type(None))):\n        raise ValueError(\"flagellar_points must be a list of Point3D objects or None\")\n    \n    norm_tomo = (tomography - tomography.min()) / (tomography.max() - tomography.min())\n    num_layers, height, width = norm_tomo.shape\n\n    z, y, x = np.mgrid[0:num_layers, 0:height, 0:width]\n\n    fig = go.Figure(data=go.Volume(\n        x=x.flatten(),\n        y=y.flatten(),\n        z=z.flatten(),\n        value=norm_tomo.flatten(),\n        isomin=0,\n        isomax=1,\n        colorscale=[[0, \"black\"], [1, \"white\"]],\n        opacityscale=[[0.0, 1.0],[0.1, 0.95],[0.2, 0.8],[0.3, 0.5],[0.4, 0.2],[1.0, 0.0]],\n        surface_count=30\n    ))\n\n    if flagellar_points is not None and len(flagellar_points) > 0:\n        colors = ['red', 'green', 'blue', 'orange', 'purple', 'cyan', 'magenta', 'yellow']\n        dark_colors = ['darkred', 'darkgreen', 'darkblue', 'darkorange', 'darkviolet', 'darkcyan', 'darkmagenta', 'gold']\n        \n        for i, point in enumerate(flagellar_points):\n            color_idx = i % len(colors)  \n            point_color = colors[color_idx]\n            dark_color = dark_colors[color_idx]\n            \n            point_type = \"Predicted\" if i == 0 else f\"Actual {i}\"\n            \n            fig.add_trace(go.Scatter3d(\n                x=[point.x],\n                y=[point.y],\n                z=[point.z],\n                mode='markers',\n                marker=dict(\n                    size=8,\n                    color=point_color,\n                    symbol='circle',\n                    opacity=1.0,\n                    line=dict(\n                        width=3,\n                        color=dark_color\n                    ),\n                    sizemode='diameter'\n                ),\n                name=f'{point_type} Motor',\n                text=[f'{point_type} Motor: ({point.z:.0f}, {point.y:.0f}, {point.x:.0f})'],\n                hovertemplate='<b>%{text}</b><br>' +\n                             'X: %{x}<br>' +\n                             'Y: %{y}<br>' +\n                             'Z: %{z}<br>' +\n                             '<extra></extra>'\n            ))\n            \n            fig.add_trace(go.Scatter3d(\n                x=[point.x],\n                y=[point.y],\n                z=[point.z],\n                mode='markers',\n                marker=dict(\n                    size=12,\n                    color=f'rgba({255 if point_color == \"red\" else 0}, {255 if point_color == \"green\" else 0}, {255 if point_color == \"blue\" else 0}, 0.3)',\n                    symbol='circle',\n                    line=dict(width=0)\n                ),\n                name=f'{point_type} Glow',\n                showlegend=False,\n                hoverinfo='skip'\n            ))\n\n    fig.update_layout(\n        title=title,\n        scene=dict(\n            camera=dict(\n                eye=dict(x=1.5, y=1.5, z=1.5)\n            )\n        )\n    )\n    \n    pio.renderers.default = \"iframe_connected\" if ON_KAGGLE else \"notebook\"\n    \n    fig.show()\n    \n\ndef visualize_cube_coords(cube_coords, tomography_shape: BoxShape, title='title not set', offset=(0, 0, 0)):\n\n    def plot_cube(ax, cube_coord: Box3D , color='r', offset=(0, 0, 0)):\n\n        if not isinstance(cube_coord, Box3D):\n            raise ValueError(\"cube_coord must be a Box2D object\")\n        \n        z0 = cube_coord.z_start\n        z1 = cube_coord.z_end\n        y0 = cube_coord.y_start\n        y1 = cube_coord.y_end\n        x0 = cube_coord.x_start\n        x1 = cube_coord.x_end\n\n        z_off, y_off, x_off = offset\n        \n        z0, z1 = z0 + z_off, z1 + z_off\n        y0, y1 = y0 + y_off, y1 + y_off\n        x0, x1 = x0 + x_off, x1 + x_off\n        \n        vertices = np.array([\n            [x0, y0, z0],\n            [x1, y0, z0],\n            [x1, y1, z0],\n            [x0, y1, z0],\n            [x0, y0, z1],\n            [x1, y0, z1],\n            [x1, y1, z1],\n            [x0, y1, z1],\n        ])\n\n        faces = [\n            [vertices[0], vertices[1], vertices[2], vertices[3]],\n            [vertices[4], vertices[5], vertices[6], vertices[7]],\n            [vertices[0], vertices[1], vertices[5], vertices[4]],\n            [vertices[2], vertices[3], vertices[7], vertices[6]],\n            [vertices[1], vertices[2], vertices[6], vertices[5]],\n            [vertices[4], vertices[7], vertices[3], vertices[0]]\n        ]\n        box = Poly3DCollection(faces, facecolors=color, edgecolors='k', alpha=0.05)\n        ax.add_collection3d(box)\n\n    fig = plt.figure(figsize=(8, 8))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.set_xlabel(r'\\(X\\)')\n    ax.set_ylabel(r'\\(Y\\)')\n    ax.set_ylabel(r'\\(Z\\)')\n\n    colors = ['red', 'blue', 'green', 'orange', 'purple']\n    for i, (cube, _, _, _) in enumerate(cube_coords):\n        plot_cube(ax, cube, color=colors[i % len(colors)], offset=offset)\n\n    ax.set_xlim(0, tomography_shape.x)\n    ax.set_ylim(0, tomography_shape.y)\n    ax.set_zlim(0, tomography_shape.z) # type: ignore\n    ax.set_title(title)\n\n    ax.set_box_aspect([float(tomography_shape.x), float(tomography_shape.y), float(tomography_shape.z)])  # type: ignore \n    \n    plt.show()\n\n\ntomo_id = \"tomo_00e047\"\noriginal_tomography_shape = getTomographyShape(train_df, tomo_id)\n\nflagellar_point = getFlagellarCoordinates(train_df, tomo_id)\ncube_coords = getCubeCoordsToExtract(original_tomography_shape, flagellar_point, minSizeRatio=FIRST_MINSIZE_RATIO, padding=EXTRACTION_PADDING)\n\ntrue_cube_coords = [cube_coord for cube_coord in cube_coords if cube_coord[1]][0]\ntrue_cube_coord_ranges = true_cube_coords[0] \nz_start, z_end = true_cube_coord_ranges.z_start, true_cube_coord_ranges.z_end\ny_start, y_end = true_cube_coord_ranges.y_start, true_cube_coord_ranges.y_end\nx_start, x_end = true_cube_coord_ranges.x_start, true_cube_coord_ranges.x_end\n\ntrue_cube_offset = (z_start, y_start, x_start)\n\ntrue_cube_shape = (z_end - z_start, y_end - y_start, x_end - x_start)\nprint(f\"True cube shape: {true_cube_shape}\")\nprint(f\"True cube coordinates: Z({z_start}-{z_end}), Y({y_start}-{y_end}), X({x_start}-{x_end})\")\n\nsub_cube_coords = getCubeCoordsToExtract(\n    original_tomography_shape=original_tomography_shape,\n    dots=flagellar_point,  \n    minSizeRatio=MINSIZE_RATIO,\n    padding=EXTRACTION_PADDING,\n    export_zone=true_cube_coord_ranges\n)\n\nvisualize_cube_coords(cube_coords, original_tomography_shape, title='3D Visualization of all the cubes that will be exported within Tomography Shape')\n\nvisualize_cube_coords([true_cube_coords], original_tomography_shape, title='3D Visualization of cube containing the flagellar within Tomography Shape')\n\nvisualize_cube_coords(sub_cube_coords, original_tomography_shape, title='3D Visualization of sub-cubes positioned within the cube containing the flagellar', offset=true_cube_offset)\n\nprint(\"Generating exports with antialiasing disabled...\")\nexports_without_aa = getTomographyExportsByTomoId(tomo_id, base_dir=base_dir, device=DEVICE, cube_coords=[true_cube_coords], output_block_shape=EXPORT_SHAPE, antialiasing=False)\n\ndisplay_tomography(exports_without_aa[0][0], title=\"Tomography cube export WITHOUT Antialiasing\", flagellar_points=exports_without_aa[0][2])   \n","metadata":{"ExecuteTime":{"end_time":"2025-05-26T13:12:19.74138Z","start_time":"2025-05-26T13:12:13.140629Z"}},"outputs":[],"execution_count":null},{"id":"0cacf3db","cell_type":"code","source":"# Displaying the same tomography but with antialiasing enabled\n\nprint(\"Generating exports with antialiasing enabled...\")\nexports_with_aa = getTomographyExportsByTomoId(tomo_id, base_dir=base_dir, device=DEVICE, cube_coords=[true_cube_coords], output_block_shape=EXPORT_SHAPE, antialiasing=True)\n\ndisplay_tomography(exports_with_aa[0][0], title=\"Tomography cube export WITH Antialiasing\", flagellar_points=exports_with_aa[0][2])\n\n# I needed a new cell for this because showing two tomographies in the same cell is not working","metadata":{},"outputs":[],"execution_count":null},{"id":"b8d9ec5a","cell_type":"markdown","source":"## 2 Loading the tomographies\nIn this section we will load the tomographies, the dataset and the labels. ","metadata":{}},{"id":"37f62017d6a0afef","cell_type":"markdown","source":"### 2.1 Multi-GPU Data Export with Caching\nHere we will export the tomographies into multiple blocks, and save them in a cache directory if they are not already present. We will not use the antialiasing as from some testing it seems that the model trains better without it.","metadata":{}},{"id":"a73d38052b0af11a","cell_type":"code","source":"tomo_ids = train_df['tomo_id'].unique().tolist()\n\nif DEBUG:\n    num_debug_tomos = min(20, len(tomo_ids))\n    print(f\"--- DEBUG MODE ENABLED: Processing only the first {num_debug_tomos} tomographies ---\")\n    tomo_ids = tomo_ids[:num_debug_tomos]\n\nprint(f\"Total tomographies to process: {len(tomo_ids)}\")\n\ncache_messages = []\ncache_messages_lock = threading.Lock()\n\ndef process_tomography(tomo_id) -> List[Tuple[np.ndarray, bool, List[Point3D], Box3D, Box3D]]:\n    \"\"\" Processes a single tomography, selecting GPU in round-robin if available \"\"\"\n    pid = os.getpid()\n    tid = threading.get_ident()\n\n    target_device = DEFAULT_DEVICE\n    acquired_semaphore = None\n    \n    if gpu_semaphores:\n        for device_id in device_ids:\n            semaphore = gpu_semaphores[device_id]\n            if semaphore.acquire(blocking=False):\n                target_device = device_id\n                acquired_semaphore = semaphore\n                break\n        \n        if acquired_semaphore is None:\n            for device_id in device_ids:\n                semaphore = gpu_semaphores[device_id]\n                try:\n                    semaphore.acquire(timeout=0.1)\n                    target_device = device_id\n                    acquired_semaphore = semaphore\n                    break\n                except:\n                    continue\n            \n            if acquired_semaphore is None:\n                target_device = device_ids[0]\n                acquired_semaphore = gpu_semaphores[target_device]\n                acquired_semaphore.acquire()\n\n    if target_device.startswith('cuda'):\n        torch.cuda.set_device(target_device)\n\n    try:\n        tomography_shape = getTomographyShape(train_df, tomo_id)\n        flagellar_points = getFlagellarCoordinates(train_df, tomo_id)\n        cube_coords = getCubeCoordsToExtract(tomography_shape, flagellar_points, FIRST_MINSIZE_RATIO, padding=EXTRACTION_PADDING)\n\n        tomo_hash = hashlib.md5(str(tomo_id).encode()).hexdigest()[:10]\n        cube_hash = hashlib.md5(str([(c[0], c[1]) for c in cube_coords]).encode()).hexdigest()[:10]\n        padding_hash = hashlib.md5(str(EXTRACTION_PADDING).encode()).hexdigest()[:6]\n        export_shape_hash = hashlib.md5(str(EXPORT_SHAPE.toTuple()).encode()).hexdigest()[:6]\n        iterations_hash = hashlib.md5(str(EXTRACTION_ITERATIONS).encode()).hexdigest()[:6]  \n        minSizeRatio_hash = hashlib.md5(str(FIRST_MINSIZE_RATIO).encode()).hexdigest()[:6]\n        firstMinSizeRatio_hash = hashlib.md5(str(FIRST_MINSIZE_RATIO).encode()).hexdigest()[:6]\n        aa_flag = \"aa\" if ANTIALIASING else \"noaa\"\n\n        cache_filename = f\"exports_{tomo_hash}_{export_shape_hash}_{cube_hash}_{padding_hash}_{iterations_hash}_{aa_flag}_{minSizeRatio_hash}_{firstMinSizeRatio_hash}.pkl\"\n        cache_path = os.path.join(\"cache\", cache_filename)\n        \n        os.makedirs(\"cache\", exist_ok=True)\n\n        if target_device.startswith('cuda'):\n            with torch.cuda.device(target_device):\n                \n                if os.path.exists(cache_path):\n                    with cache_messages_lock:\n                        cache_messages.append(f\"[{pid}-{tid}] Loading cached results for {tomo_id} from {cache_filename} on {target_device}\")\n                    with open(cache_path, 'rb') as f:\n                        tomography_export = pickle.load(f)\n                else:\n                    with cache_messages_lock:\n                        cache_messages.append(f\"[{pid}-{tid}] No cache found for {tomo_id}. Generating exports on {target_device}...\")\n                    tomography_export = getTomographyExportsByTomoId(\n                        tomo_id,\n                        device=target_device,\n                        base_dir=base_dir,\n                        cube_coords=cube_coords,\n                        output_block_shape=EXPORT_SHAPE,\n                        antialiasing=ANTIALIASING,\n                        iterations=EXTRACTION_ITERATIONS\n                    )\n                    \n        else:\n            if os.path.exists(cache_path):\n                with cache_messages_lock:\n                    cache_messages.append(f\"[{pid}-{tid}] Loading cached results for {tomo_id} from {cache_filename} on {target_device}\")\n                with open(cache_path, 'rb') as f:\n                    tomography_export = pickle.load(f)\n            else:\n                with cache_messages_lock:\n                    cache_messages.append(f\"[{pid}-{tid}] No cache found for {tomo_id}. Generating exports on {target_device}...\")\n                tomography_export = getTomographyExportsByTomoId(\n                    tomo_id,\n                    device=target_device,\n                    base_dir=base_dir,\n                    cube_coords=cube_coords,\n                    output_block_shape=EXPORT_SHAPE,\n                    antialiasing=ANTIALIASING,\n                    iterations=EXTRACTION_ITERATIONS\n                )\n                \n        with open(cache_path, 'wb') as f:\n            pickle.dump(tomography_export, f)\n        with cache_messages_lock:\n            cache_messages.append(f\"[{pid}-{tid}] Cached exports for {tomo_id} in {cache_filename} (processed on {target_device})\")\n        return tomography_export\n    except Exception as e:\n        print(f\"[{pid}-{tid}] !!! ERROR processing {tomo_id} on device {target_device}: {e}\\n--- TRACEBACK ---\\n{traceback.format_exc()}\\n--- END TRACEBACK ---\")\n        return []\n    finally:\n        if acquired_semaphore:\n            acquired_semaphore.release()\n\n\nExecutor = concurrent.futures.ThreadPoolExecutor\n\nexports = []\nresults = []\n\n# --- Execute using Threads ---\ntry:\n    with Executor(max_workers=EXPORT_WORKERS) as executor:\n        results_iterator = executor.map(process_tomography, tomo_ids)\n        results = list(tqdm(results_iterator,\n                            total=len(tomo_ids),\n                            desc=f\"Processing Tomographies (ThreadPoolExecutor - {len(device_ids)} GPU(s))\"))\nexcept Exception as e:\n    print(f\"!!! ERROR during ThreadPoolExecutor execution: {e}\")\n    print(traceback.format_exc())\n    results = []\n\n# --- Print cache messages after tqdm completes ---\nprint(\"\\n--- Cache Status Summary ---\")\nif DEBUG:\n    for message in cache_messages:\n        print(message)\n    print(f\"Total cache messages: {len(cache_messages)}\")\n    print(\"--- End Cache Summary ---\\n\")\n\nprint(\"Collecting results...\")\nfor not_none in results:\n    if not_none:\n        exports.extend(not_none)\n\ntotal_size = 0\nfor export_data, _, _, _, _ in exports:\n    if isinstance(export_data, np.ndarray):\n        total_size += export_data.nbytes\n    elif isinstance(export_data, torch.Tensor):\n         total_size += export_data.nelement() * export_data.element_size()\n\nif exports:\n    random.shuffle(exports)\n    print(\"Exports shuffled.\")\nelse:\n    print(\"Warning: No exports were generated or collected.\")\n\nprint(f\"Number of exports successfully collected: {len(exports)}\")\nprint(f\"Total estimated size of exports in memory: {total_size / (1024**2):.2f} MB\")","metadata":{},"outputs":[],"execution_count":null},{"id":"a6ea76df","cell_type":"markdown","source":"### 2.3 Dataset Preparation\nThe dataset for training will present random enanchements to the blocks, and will be split in train and validation sets. The dataset for testing will not present any augmentation, and will be used to evaluate the model performance.","metadata":{}},{"id":"3471fc8f40ff6ac","cell_type":"code","source":"class TomographyDataset(Dataset):\n   def __init__(self, exports, augment=True, noise_std=0.02, mirror_prob=0.5,\n                elastic_prob=0.15, zoom_out_prob=0.2, zoom_out_range=(0.85, 0.95),\n                brightness_prob=0.2, brightness_range=0.05, contrast_prob=0.2, contrast_range=0.05,\n                blur_prob=0.2, crop_shift_prob=0.3):\n       \n       self.augment = augment\n       self.noise_std = noise_std\n       self.mirror_prob = mirror_prob\n       \n       self.elastic_prob = elastic_prob\n       self.zoom_out_prob = zoom_out_prob\n       self.zoom_out_range = zoom_out_range\n       self.brightness_prob = brightness_prob\n       self.brightness_range = brightness_range\n       self.contrast_prob = contrast_prob\n       self.contrast_range = contrast_range\n       self.blur_prob = blur_prob\n       self.crop_shift_prob = crop_shift_prob\n       \n       self.data = []\n       \n       for volume, label, _, _, _ in exports:\n           self.data.append((volume, 1 if label else 0))\n           \n       positive_count = sum(1 for _, label in self.data if label == 1)\n       negative_count = len(self.data) - positive_count\n       print(f\"Dataset created with {positive_count} positive and {negative_count} negative samples\")\n\n   def __len__(self):\n       return len(self.data)\n   \n   def add_noise(self, volume):\n       \"\"\"Add Gaussian noise to the volume\"\"\"\n       if self.noise_std > 0:\n           noise = torch.randn_like(volume) * self.noise_std\n           volume = volume + noise\n           volume = torch.clamp(volume, 0.0, 1.0)\n       return volume\n   \n   def apply_mirroring(self, volume):\n       \"\"\"Apply random mirroring across the 3 spatial dimensions\"\"\"        \n       if random.random() < self.mirror_prob:\n           volume = torch.flip(volume, dims=[1])\n       \n       if random.random() < self.mirror_prob:\n           volume = torch.flip(volume, dims=[2])\n           \n       if random.random() < self.mirror_prob:\n           volume = torch.flip(volume, dims=[3])\n           \n       return volume\n\n   def apply_zoom_out(self, volume):\n       \"\"\"Apply random zoom out (making cube smaller) only\"\"\"\n       if random.random() < self.zoom_out_prob:\n           zoom_factor = random.uniform(self.zoom_out_range[0], self.zoom_out_range[1])\n           \n           D, H, W = volume.shape[1:]\n           new_size = (int(D * zoom_factor), int(H * zoom_factor), int(W * zoom_factor))\n           \n           volume_resized = F.interpolate(volume.unsqueeze(0), size=new_size, \n                                        mode='trilinear', align_corners=False).squeeze(0)\n           \n           volume = self._center_pad(volume_resized, (D, H, W))\n           \n       return volume\n\n   def _center_pad(self, volume, target_size):\n       \"\"\"Center pad volume to target size\"\"\"\n       current_size = volume.shape[1:]\n       target_D, target_H, target_W = target_size\n       \n       def get_pad(current, target):\n           pad_total = target - current\n           pad_before = pad_total // 2\n           pad_after = pad_total - pad_before\n           return (pad_before, pad_after)\n       \n       d_pad = get_pad(current_size[0], target_D)\n       h_pad = get_pad(current_size[1], target_H)\n       w_pad = get_pad(current_size[2], target_W)\n       \n       padding = (w_pad[0], w_pad[1], h_pad[0], h_pad[1], d_pad[0], d_pad[1])\n       volume = F.pad(volume, padding, mode='reflect')\n       \n       return volume\n\n   def apply_elastic_deformation_simple(self, volume):\n       \"\"\"Simple elastic deformation using small random shifts\"\"\"\n       if random.random() < self.elastic_prob:\n           D, H, W = volume.shape[1:]\n           \n           max_displacement = 1.5\n           \n           num_control_points = random.randint(2, 4)\n           \n           for _ in range(num_control_points):\n               center_d = random.randint(2, D-3)\n               center_h = random.randint(2, H-3) \n               center_w = random.randint(2, W-3)\n               \n               shift_d = random.uniform(-max_displacement, max_displacement)\n               shift_h = random.uniform(-max_displacement, max_displacement)\n               shift_w = random.uniform(-max_displacement, max_displacement)\n               \n               for dd in range(-1, 2):\n                   for dh in range(-1, 2):\n                       for dw in range(-1, 2):\n                           d_idx = center_d + dd\n                           h_idx = center_h + dh  \n                           w_idx = center_w + dw\n                           \n                           if 0 <= d_idx < D and 0 <= h_idx < H and 0 <= w_idx < W:\n                               src_d = max(0, min(D-1, int(d_idx + shift_d)))\n                               src_h = max(0, min(H-1, int(h_idx + shift_h)))\n                               src_w = max(0, min(W-1, int(w_idx + shift_w)))\n                               \n                               blend_factor = 0.2\n                               original = volume[0, d_idx, h_idx, w_idx]\n                               shifted = volume[0, src_d, src_h, src_w]\n                               volume[0, d_idx, h_idx, w_idx] = (1 - blend_factor) * original + blend_factor * shifted\n       \n       return volume\n\n   def apply_gaussian_blur(self, volume):\n       \"\"\"Apply slight Gaussian blur to simulate imaging artifacts\"\"\"\n       if random.random() < self.blur_prob:\n           kernel_size = 3\n           sigma = random.uniform(0.3, 0.6)\n           \n           ax = torch.arange(-kernel_size // 2 + 1., kernel_size // 2 + 1.)\n           xx, yy, zz = torch.meshgrid(ax, ax, ax, indexing='ij')\n           kernel = torch.exp(-(xx**2 + yy**2 + zz**2) / (2 * sigma**2))\n           kernel = kernel / kernel.sum()\n           kernel = kernel.unsqueeze(0).unsqueeze(0)\n           \n           volume_padded = F.pad(volume.unsqueeze(0), \n                                (kernel_size//2, kernel_size//2, \n                                 kernel_size//2, kernel_size//2,\n                                 kernel_size//2, kernel_size//2), \n                                mode='reflect')\n           \n           blurred = F.conv3d(volume_padded, kernel, padding=0)\n           volume = blurred.squeeze(0)\n       \n       return volume\n\n   def apply_random_crop_and_shift(self, volume):\n       \"\"\"Apply small random crops and shifts to simulate position variations\"\"\"\n       if random.random() < self.crop_shift_prob:\n           D, H, W = volume.shape[1:]\n           \n           shift_size = 1\n           \n           shift_d = random.randint(-shift_size, shift_size)\n           shift_h = random.randint(-shift_size, shift_size)\n           shift_w = random.randint(-shift_size, shift_size)\n           \n           start_d = max(0, shift_d)\n           start_h = max(0, shift_h)\n           start_w = max(0, shift_w)\n           \n           end_d = min(D, D + shift_d)\n           end_h = min(H, H + shift_h)\n           end_w = min(W, W + shift_w)\n           \n           cropped = volume[:, start_d:end_d, start_h:end_h, start_w:end_w]\n           \n           pad_d_before = max(0, -shift_d)\n           pad_d_after = max(0, shift_d)\n           pad_h_before = max(0, -shift_h)\n           pad_h_after = max(0, shift_h)\n           pad_w_before = max(0, -shift_w)\n           pad_w_after = max(0, shift_w)\n           \n           padding = (pad_w_before, pad_w_after, pad_h_before, pad_h_after, pad_d_before, pad_d_after)\n           volume = F.pad(cropped, padding, mode='reflect')\n       \n       return volume\n\n   def apply_intensity_variations(self, volume):\n       \"\"\"Apply brightness and contrast changes\"\"\"\n       if random.random() < self.brightness_prob:\n           brightness_factor = random.uniform(-self.brightness_range, self.brightness_range)\n           volume = volume + brightness_factor\n           volume = torch.clamp(volume, 0.0, 1.0)\n       \n       if random.random() < self.contrast_prob:\n           contrast_factor = random.uniform(1.0 - self.contrast_range, 1.0 + self.contrast_range)\n           mean_val = volume.mean()\n           volume = (volume - mean_val) * contrast_factor + mean_val\n           volume = torch.clamp(volume, 0.0, 1.0)\n       \n       return volume\n\n   def __getitem__(self, idx):\n       volume_np, label = self.data[idx]\n       \n       volume_norm = volume_np.astype(np.float32) / 255.0\n       volume = torch.from_numpy(volume_norm).unsqueeze(0)\n       label_tensor = torch.tensor(label, dtype=torch.long)\n       \n       if self.augment:\n           volume = volume.clone()\n           \n           volume = self.add_noise(volume)\n           volume = self.apply_mirroring(volume)\n           volume = self.apply_zoom_out(volume)\n           volume = self.apply_random_crop_and_shift(volume)\n           volume = self.apply_elastic_deformation_simple(volume)\n           volume = self.apply_gaussian_blur(volume)\n           volume = self.apply_intensity_variations(volume)\n           \n       return volume, label_tensor\n\nprint(\"Preparing dataset with enhanced augmentations (no rotation, zoom in, or dropout)...\")\n\nfrom sklearn.model_selection import train_test_split\n\ndataset_no_aug = TomographyDataset(\n   exports, \n   augment=False\n)\n\nclass_counts = [0, 0]\nfor _, label in dataset_no_aug.data:\n   class_counts[label] += 1\nprint(f\"Original class distribution - Negative: {class_counts[0]}, Positive: {class_counts[1]}\")\n\nall_indices = list(range(len(dataset_no_aug)))\nall_labels = [dataset_no_aug.data[i][1] for i in all_indices]\n\ntrain_indices, val_indices = train_test_split(\n   all_indices,\n   test_size=VALIDATION_SET_SIZE,\n   stratify=all_labels,\n   random_state=SEED\n)\n\nprint(f\"Training samples: {len(train_indices)}\")\nprint(f\"Validation samples: {len(val_indices)}\")\n\ntrain_labels = [all_labels[i] for i in train_indices]\nval_labels = [all_labels[i] for i in val_indices]\n\nprint(f\"Train set - Negative: {train_labels.count(0)}, Positive: {train_labels.count(1)}\")\nprint(f\"Val set - Negative: {val_labels.count(0)}, Positive: {val_labels.count(1)}\")\n\ntrain_exports = [exports[i] for i in train_indices]\nval_exports = [exports[i] for i in val_indices]\n\ntrain_dataset = TomographyDataset(\n   train_exports,\n   augment=True,\n   noise_std=TRAINING_NOISE,\n   mirror_prob=TRAINING_MIRROR_PROB,\n   elastic_prob=0.1,\n   zoom_out_prob=0.1,\n   zoom_out_range=(0.85, 0.95),\n   brightness_prob=0.1,\n   brightness_range=0.05,\n   contrast_prob=0.1,\n   contrast_range=0.05,\n   blur_prob=0.1,\n   crop_shift_prob=0.1\n)\n\nval_dataset = TomographyDataset(\n   val_exports,\n   augment=False\n)\n\ntrain_loader = DataLoader(\n   train_dataset,\n   batch_size=BATCH_SIZE,\n   shuffle=True,\n   num_workers=4 if not sys.platform == \"darwin\" else 0\n)\n\nval_loader = DataLoader(\n   val_dataset,\n   batch_size=BATCH_SIZE,\n   shuffle=False,\n   num_workers=4 if not sys.platform == \"darwin\" else 0\n)\n\nprint(\"Dataset preparation completed with stratified split and appropriate augmentation.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"ea4faeb1","cell_type":"markdown","source":"## 3 Neural Network and Training","metadata":{}},{"id":"5b1eda001f68524f","cell_type":"markdown","source":"### 3.1 3D ResNet-Style CNN Architecture\nHere we implement a 3D CNN to classify tomography exports as containing flagellar proteins or not.\n","metadata":{}},{"id":"1f9df270","cell_type":"code","source":"class FlagellarClassifier3D(nn.Module):\n   def __init__(self, input_shape):\n       super(FlagellarClassifier3D, self).__init__()\n       \n       self.depth, self.height, self.width = input_shape\n       \n       self.conv1 = nn.Conv3d(1, 32, kernel_size=3, padding=1)\n       self.bn1 = nn.BatchNorm3d(32)\n       \n       self.conv2 = nn.Conv3d(32, 64, kernel_size=3, padding=1)\n       self.bn2 = nn.BatchNorm3d(64)\n       \n       self.conv3 = nn.Conv3d(64, 128, kernel_size=3, padding=1)\n       self.bn3 = nn.BatchNorm3d(128)\n       \n       self.se_block = nn.Sequential(\n           nn.AdaptiveAvgPool3d(1),\n           nn.Conv3d(128, 32, kernel_size=1),\n           nn.ReLU(),\n           nn.Conv3d(32, 128, kernel_size=1),\n           nn.Sigmoid()\n       )\n       \n       self.conv4 = nn.Conv3d(128, 256, kernel_size=3, padding=1)\n       self.bn4 = nn.BatchNorm3d(256)\n       \n       self.global_avg_pool = nn.AdaptiveAvgPool3d((2, 2, 2))\n       \n       self.classifier = nn.Sequential(\n           nn.Linear(256 * 8, 256),\n           nn.ReLU(),\n           nn.Dropout(DROPOUT_RATE),\n           nn.Linear(256, 64),\n           nn.ReLU(),\n           nn.Dropout(DROPOUT_RATE * 0.5),\n           nn.Linear(64, 2)\n       )\n       \n   def forward(self, x):\n       x = F.relu(self.bn1(self.conv1(x)))\n       x = F.max_pool3d(x, 2)\n       \n       x = F.relu(self.bn2(self.conv2(x)))\n       x = F.max_pool3d(x, 2)\n       \n       x = F.relu(self.bn3(self.conv3(x)))\n       \n       se_weights = self.se_block(x)\n       x = x * se_weights\n       \n       x = F.relu(self.bn4(self.conv4(x)))\n       \n       x = self.global_avg_pool(x)\n       x = x.view(x.size(0), -1)\n       \n       x = self.classifier(x)\n       return x","metadata":{},"outputs":[],"execution_count":null},{"id":"70a2a6e4","cell_type":"markdown","source":"### 3.2 Model Initialization and Loss Configuration","metadata":{}},{"id":"9c967a3f","cell_type":"code","source":"model = FlagellarClassifier3D(EXPORT_SHAPE.toTuple()).to(DEVICE)\n\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs for training!\")\n    model = nn.DataParallel(model)\n\nprint(model)\n\ncounts = [0,0]\nfor e in exports:\n    if e[1]: counts[1] += 1\n    else: counts[0] += 1\n\nprint(f\"Class weights: Negative: {counts[0]}, Positive: {counts[1]}\")\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n    \n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none')\n        pt = torch.exp(-ce_loss)\n        \n        if self.alpha is not None:\n            if self.alpha.device != targets.device:\n                self.alpha = self.alpha.to(targets.device)\n            alpha_t = self.alpha[targets]\n            focal_loss = alpha_t * (1 - pt) ** self.gamma * ce_loss\n        else:\n            focal_loss = (1 - pt) ** self.gamma * ce_loss\n        \n        if self.reduction == 'mean':\n            return focal_loss.mean()\n        elif self.reduction == 'sum':\n            return focal_loss.sum()\n        else:\n            return focal_loss\n\nimbalance_ratio = counts[0] / counts[1] if counts[1] > 0 else 1.0\nalpha_positive = imbalance_ratio * ALPHA_RATIO\nalpha_negative = 1.0\n\nalpha_tensor = torch.tensor([alpha_negative, alpha_positive], dtype=torch.float32, device=DEVICE)\n\ncriterion = FocalLoss(alpha=alpha_tensor, gamma=1.2, reduction='mean')\nprint(f\"Using Focal Loss with alpha=[{alpha_negative:.2f}, {alpha_positive:.2f}], gamma=2.0\")\nprint(f\"Class imbalance ratio: {imbalance_ratio:.2f}\")\n\n\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, \n    mode='max',\n    factor=0.5,\n    patience=8,\n    threshold=0.002,\n    min_lr=1e-5\n    )","metadata":{},"outputs":[],"execution_count":null},{"id":"5b611dc8","cell_type":"markdown","source":"### 3.3 Evaluation Function","metadata":{}},{"id":"22868c46","cell_type":"code","source":"\ndef evaluate(model, dataloader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for inputs, labels in dataloader:\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs, 1)\n\n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    val_loss = running_loss / len(dataloader.dataset)\n    accuracy = accuracy_score(all_labels, all_preds)\n    precision = precision_score(all_labels, all_preds, zero_division=0)\n    recall = recall_score(all_labels, all_preds, zero_division=0)\n    f1 = f1_score(all_labels, all_preds, zero_division=0)\n    \n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    \n    # Negative class accuracy (True Negative Rate / Specificity)\n    negative_mask = all_labels == 0\n    if np.sum(negative_mask) > 0:\n        negative_accuracy = np.sum((all_preds == 0) & negative_mask) / np.sum(negative_mask)\n    else:\n        negative_accuracy = 0.0\n    \n    # Positive class accuracy (True Positive Rate / Sensitivity / Recall)\n    positive_mask = all_labels == 1\n    if np.sum(positive_mask) > 0:\n        positive_accuracy = np.sum((all_preds == 1) & positive_mask) / np.sum(positive_mask)\n    else:\n        positive_accuracy = 0.0\n\n    metrics = {\n        'loss': val_loss,\n        'accuracy': accuracy,\n        'precision': precision,\n        'recall': recall,\n        'f1_score': f1,\n        'negative_accuracy': negative_accuracy,\n        'positive_accuracy': positive_accuracy\n    }\n\n    return metrics","metadata":{},"outputs":[],"execution_count":null},{"id":"85e91605","cell_type":"markdown","source":"### 3.4 Training Loop","metadata":{}},{"id":"bb4b1399ffe70ffa","cell_type":"code","source":"print(f\"Starting training on {DEVICE}...\")\nbest_f1 = -1\nhistory = {\n    'train_loss': [],\n    'train_metrics': [],\n    'val_metrics': [],\n    'learning_rates': []\n}\n\nif DEBUG:\n    EPOCHS = 1\n\nfor epoch in range(EPOCHS): # type: ignore\n    print(f\"Epoch {epoch+1}/{EPOCHS}\") # type: ignore\n    \n    model.train()\n    running_loss = 0.0\n    all_train_preds = []\n    all_train_labels = []\n    \n    progress_bar = tqdm(train_loader, desc=f\"Training\", leave=True)\n    for batch_idx, (inputs, labels) in enumerate(progress_bar):\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * inputs.size(0)\n        _, predicted = torch.max(outputs, 1)\n        \n        all_train_preds.extend(predicted.cpu().numpy())\n        all_train_labels.extend(labels.cpu().numpy())\n\n        current_loss = running_loss / len(all_train_labels)\n        current_acc = sum(p == l for p, l in zip(all_train_preds, all_train_labels)) / len(all_train_labels)\n        progress_bar.set_postfix({\n            'loss': f\"{current_loss:.4f}\",\n            'acc': f\"{current_acc:.4f}\"\n        })\n\n    train_loss = running_loss / len(all_train_labels)\n    train_accuracy = accuracy_score(all_train_labels, all_train_preds)\n    train_precision = precision_score(all_train_labels, all_train_preds, zero_division=0)\n    train_recall = recall_score(all_train_labels, all_train_preds, zero_division=0)\n    train_f1 = f1_score(all_train_labels, all_train_preds, zero_division=0)\n    \n    all_train_preds_np = np.array(all_train_preds)\n    all_train_labels_np = np.array(all_train_labels)\n    \n    negative_mask = all_train_labels_np == 0\n    if np.sum(negative_mask) > 0:\n        train_negative_accuracy = np.sum((all_train_preds_np == 0) & negative_mask) / np.sum(negative_mask)\n    else:\n        train_negative_accuracy = 0.0\n    \n    positive_mask = all_train_labels_np == 1\n    if np.sum(positive_mask) > 0:\n        train_positive_accuracy = np.sum((all_train_preds_np == 1) & positive_mask) / np.sum(positive_mask)\n    else:\n        train_positive_accuracy = 0.0\n    \n    train_metrics = {\n        'loss': train_loss,\n        'accuracy': train_accuracy,\n        'precision': train_precision,\n        'recall': train_recall,\n        'f1_score': train_f1,\n        'negative_accuracy': train_negative_accuracy,\n        'positive_accuracy': train_positive_accuracy\n    }\n\n    history['train_loss'].append(train_loss)\n    history['train_metrics'].append(train_metrics)\n    history['learning_rates'].append(optimizer.param_groups[0]['lr'])\n    \n    val_metrics = evaluate(model, val_loader, criterion, DEVICE)\n    history['val_metrics'].append(val_metrics)\n\n    scheduler.step(val_metrics['f1_score'])\n    \n    print(f\"     Epoch {epoch+1}:         LR: {optimizer.param_groups[0]['lr']:.7f}\")\n    print(f\"     Train Loss: {train_loss:.4f},       Train F1: {train_f1:.4f}\")\n    print(f\"       Val Loss: {val_metrics['loss']:.4f},         Val F1: {val_metrics['f1_score']:.4f}\")\n    print(f\"Train Precision: {train_precision:.4f},   Train Recall: {train_recall:.4f}\")\n    print(f\"  Val Precision: {val_metrics['precision']:.4f},     Val Recall: {val_metrics['recall']:.4f}\")\n    \n    if val_metrics['f1_score'] > best_f1:\n        best_f1 = val_metrics['f1_score']\n        torch.save(model.state_dict(), 'best_flagellar_classifier.pth')\n        print(\"Model saved!\")\n    \n    print(\"-\" * 50)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a2b85258","cell_type":"markdown","source":"### 3.5 Plot training and validation metrics","metadata":{}},{"id":"815714bc41249009","cell_type":"code","source":"plt.figure(figsize=(18, 12))\n\nplt.subplot(2, 3, 1)\nax1 = plt.gca()\nline1 = ax1.plot(history['train_loss'], label='Train Loss', linewidth=2, color='blue')\nline2 = ax1.plot([m['loss'] for m in history['val_metrics']], label='Val Loss', linewidth=2, color='orange')\nax1.set_xlabel('Epoch')\nax1.set_ylabel('Loss', color='black')\nax1.tick_params(axis='y', labelcolor='black')\nax1.grid(True, alpha=0.3)\n\nax2 = ax1.twinx()\nline3 = ax2.plot(history['learning_rates'], label='Learning Rate', linewidth=2, color='red', linestyle='--', alpha=0.7)\nax2.set_ylabel('Learning Rate', color='red')\nax2.tick_params(axis='y', labelcolor='red')\n\nlines = line1 + line2 + line3\nlabels = [str(l.get_label()) for l in lines]\nax1.legend(lines, labels, loc='upper right')\n\nplt.title('Loss Metrics & Learning Rate')\n\nplt.subplot(2, 3, 2)\nplt.plot([m['f1_score'] for m in history['train_metrics']], label='Train F1', linewidth=2)\nplt.plot([m['f1_score'] for m in history['val_metrics']], label='Val F1', linewidth=2)\nplt.title('F1 Score')\nplt.xlabel('Epoch')\nplt.ylabel('F1 Score')\nplt.legend()\nplt.grid(True, alpha=0.3)\n\nplt.subplot(2, 3, 3)\nplt.plot([m['precision'] for m in history['train_metrics']], label='Train Precision', linewidth=2)\nplt.plot([m['precision'] for m in history['val_metrics']], label='Val Precision', linewidth=2)\nplt.title('Precision')\nplt.xlabel('Epoch')\nplt.ylabel('Precision')\nplt.legend()\nplt.grid(True, alpha=0.3)\n\nplt.subplot(2, 3, 4)\nplt.plot([m['recall'] for m in history['train_metrics']], label='Train Recall', linewidth=2)\nplt.plot([m['recall'] for m in history['val_metrics']], label='Val Recall', linewidth=2)\nplt.title('Recall')\nplt.xlabel('Epoch')\nplt.ylabel('Recall')\nplt.legend()\nplt.grid(True, alpha=0.3)\n\nplt.subplot(2, 3, 5)\nplt.plot([m.get('negative_accuracy', 0) for m in history['train_metrics']], label='Train Negative Accuracy', linewidth=2, color='blue')\nplt.plot([m.get('negative_accuracy', 0) for m in history['val_metrics']], label='Val Negative Accuracy', linewidth=2, color='orange')\nplt.title('Negative Class Accuracy\\n(% of negatives guessed right)')\nplt.xlabel('Epoch')\nplt.ylabel('Negative Accuracy')\nplt.legend()\nplt.grid(True, alpha=0.3)\n\n\nplt.subplot(2, 3, 6)\nplt.plot([m.get('positive_accuracy', 0) for m in history['train_metrics']], label='Train Positive Accuracy', linewidth=2, color='blue')\nplt.plot([m.get('positive_accuracy', 0) for m in history['val_metrics']], label='Val Positive Accuracy', linewidth=2, color='orange')\nplt.title('Positive Class Accuracy\\n(% of positives guessed right)')\nplt.xlabel('Epoch')\nplt.ylabel('Positive Accuracy')\nplt.legend()\nplt.grid(True, alpha=0.3)\n\n\nhyperparams_text = f\"\"\"Hyperparameters: Batch Size: {BATCH_SIZE}, Epochs: {EPOCHS}\nLearning Rate: {LR}, Gamma: {GAMMA}, Val Split: {VALIDATION_SET_SIZE}, Dropout Rate: {DROPOUT_RATE}\nExport Shape: {EXPORT_SHAPE}, Extraction Iter: {EXTRACTION_ITERATIONS}, Padding: {EXTRACTION_PADDING}\nAntialiasing: {ANTIALIASING}, Training Noise: {TRAINING_NOISE}, Training Mirror Prob: {TRAINING_MIRROR_PROB}\"\"\"\n\nplt.figtext(0.02, 0.02, hyperparams_text, fontsize=9, \n           bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8),\n           verticalalignment='bottom')\n\nplt.tight_layout()\nplt.subplots_adjust(bottom=0.25)\n\noutput_dir = \"training_plots\"\nos.makedirs(output_dir, exist_ok=True)\n\nshape_str = f\"{EXPORT_SHAPE.z}x{EXPORT_SHAPE.y}x{EXPORT_SHAPE.x}\"\naa_str = \"aa\" if ANTIALIASING else \"noaa\"\nfilename = f\"training_metrics_bs{BATCH_SIZE}_ep{EPOCHS}_lr{LR}_gamma{GAMMA}_shape{shape_str}_iter{EXTRACTION_ITERATIONS}_pad{EXTRACTION_PADDING}_{aa_str}_{TRAINING_NOISE}_{TRAINING_MIRROR_PROB}_{DROPOUT_RATE}.png\"\nfull_path = os.path.join(output_dir, filename)\n\nplt.savefig(full_path, dpi=300, bbox_inches='tight')\nprint(f\"Figure saved to: {os.path.abspath(full_path)}\")\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"a130e9ae","cell_type":"markdown","source":"## 4 Submssion generation and validation","metadata":{}},{"id":"c8658571","cell_type":"markdown","source":"### 4.1 Model Loading and Setup\nLoading the best saved model","metadata":{}},{"id":"fe216eaede58bb6c","cell_type":"code","source":"model_path = 'best_flagellar_classifier.pth'\nloaded_model = FlagellarClassifier3D(EXPORT_SHAPE.toTuple()).to(DEVICE)\n\nif torch.cuda.device_count() > 1 and not isinstance(loaded_model, nn.DataParallel):\n    print(\"Wrapping model with nn.DataParallel for loading.\")\n    loaded_model = nn.DataParallel(loaded_model)\n\nprint(f\"Loading model weights from {model_path} to {DEVICE}...\")\nloaded_model.load_state_dict(torch.load(model_path, map_location=DEVICE, weights_only=True))\nloaded_model.eval()\nprint(f\"Model {model_path} loaded successfully and set to evaluation mode.\")\n\ntest_data_base_dir = os.path.join(base_dir, 'test')\n\ntest_tomo_ids = [d for d in os.listdir(test_data_base_dir) if os.path.isdir(os.path.join(test_data_base_dir, d))]\n\nprint(f\"Found {len(test_tomo_ids)} tomographies in the test set: {test_tomo_ids[:5]}...\")","metadata":{},"outputs":[],"execution_count":null},{"id":"61dedb56","cell_type":"markdown","source":"### 4.2 Cube Evaluation Functions\nThese function are used to evaluate the test tomographies, and provide the best blocks to be used for the submission.","metadata":{}},{"id":"01d10719","cell_type":"code","source":"def evaluate_cube_probability(resized_block, model):\n\n    volume_norm = resized_block.astype(np.float32) / 255.0\n    volume_tensor = torch.from_numpy(volume_norm).unsqueeze(0)\n    volume_tensor = volume_tensor.unsqueeze(0).to(DEVICE)\n    \n    with torch.no_grad():\n        outputs = model(volume_tensor)\n        probabilities_tensor = F.softmax(outputs, dim=1)\n        predicted_class_tensor = torch.argmax(probabilities_tensor, dim=1)\n        predicted_class = predicted_class_tensor.item()\n        \n        class1_probability = probabilities_tensor[0, 1].item()\n        probabilities = probabilities_tensor.cpu().numpy().tolist()[0]\n    \n    return {\n        'predicted_class': predicted_class,\n        'probabilities': probabilities,\n        'class1_probability': class1_probability\n    }\n\ndef find_best_cube(tomography, cube_coords_data: List[Tuple[Box3D, bool, List[Point3D], Box3D]], model, probability_threshold=0.60):\n\n    max_prob_class1 = 0.0\n    best_cube_coords = None\n    best_cube_info = None\n    best_original_coords = None\n    \n    exports = getTomographyExports(\n        tomography=tomography, \n        device=DEVICE,\n        output_block_shape=EXPORT_SHAPE,\n        cube_coords=cube_coords_data,\n        antialiasing=ANTIALIASING,\n        iterations=1\n    )\n    \n    for i, (resized_block, _, _, _, original_coords) in enumerate(tqdm(exports, desc=f\"Evaluating cubes for {tomo_id}\", leave=False)):\n        cube_coord_ranges = cube_coords_data[i][0]\n        \n        evaluation_result = evaluate_cube_probability(resized_block, model)\n        \n        current_prob_class1 = evaluation_result['class1_probability']\n        if current_prob_class1 > max_prob_class1:\n            max_prob_class1 = current_prob_class1\n            best_cube_coords = cube_coord_ranges\n            best_cube_info = evaluation_result.copy()\n            best_cube_info['cube_index'] = i + 1\n            best_original_coords = original_coords\n            \n    final_motor_x, final_motor_y, final_motor_z = -1, -1, -1\n\n    return {\n        'max_probability': max_prob_class1,\n        'best_cube_coords': best_cube_coords,\n        'best_cube_info': best_cube_info,\n        'final_motor_coordinates': (final_motor_x, final_motor_y, final_motor_z),\n        'meets_threshold': max_prob_class1 > probability_threshold,\n        'original_coords': best_original_coords\n    }\n\ndef find_best_cube_for_tomography(tomo_id, model, probability_threshold=0.60, iterations=ANALYSIS_ITERATIONS, padding=EXTRACTION_PADDING) -> Point3D:\n\n    image_files = getImagesPath(base_dir, \"test\", tomo_id)\n    tomography = load_tomography(image_files)\n\n    tomo_shape = tomography.shape\n\n    tomography_shape = BoxShape(\n        z=tomo_shape[0],\n        y=tomo_shape[1],\n        x=tomo_shape[2],\n        center=Point3D(tomo_shape[0] // 2, tomo_shape[1] // 2, tomo_shape[2] // 2)\n    )\n\n    best_result = None\n    export_zone = None\n\n    for i in range(iterations):\n        print(f\"Iteration {i+1}/{iterations} for tomography {tomo_id}...\")\n        \n        cube_coords_data = getCubeCoordsToExtract(\n            original_tomography_shape=tomography_shape,\n            dots=[],\n            minSizeRatio=MINSIZE_RATIO if export_zone is not None else FIRST_MINSIZE_RATIO,\n            padding=padding,\n            export_zone=export_zone\n        )\n\n        best_result = find_best_cube(tomography, cube_coords_data, model, probability_threshold)\n        \n        print(f\"  Best cube found: {best_result['original_coords']}, Probability: {best_result['max_probability']:.4f}\")\n    \n        if i == 0:\n            if not best_result[\"meets_threshold\"]:\n                print(f\"  Score {best_result['max_probability']:.4f} below threshold {probability_threshold:.4f}, exiting iterations early.\")\n                return Point3D(-1, -1, -1)\n            \n        export_zone = best_result['original_coords']\n\n    return best_result[\"original_coords\"].get_center()  # type: ignore","metadata":{},"outputs":[],"execution_count":null},{"id":"3a26aa7d","cell_type":"markdown","source":"### 4.3 Generating Submission data","metadata":{}},{"id":"6a755df0","cell_type":"code","source":"print(\"Starting submission generation...\")\n\nall_predictions_log = []\nsubmission_results = []\n\nPROBABILITY_TRESHOLD = 0.5\n\nfor tomo_id in test_tomo_ids:\n    print(f\"\\nProcessing {tomo_id}...\")\n\n    best = find_best_cube_for_tomography(\n        tomo_id, \n        loaded_model,\n        probability_threshold=PROBABILITY_TRESHOLD,\n        iterations=ANALYSIS_ITERATIONS,\n        padding=EXTRACTION_PADDING\n    ) \n    \n    # best cube in format like: ((0, 250), (470, 720), (1347, 1597))\n    \n    print (\"Best at coordinates:\", best)\n    \n    submission_results.append({\n        'tomo_id': tomo_id,\n        'Motor axis 0': best.z,\n        'Motor axis 1': best.y,\n        'Motor axis 2': best.x\n    })\n        \nprint(f\"\\nProcessed {len(submission_results)} tomographies.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"7dfe0e97","cell_type":"markdown","source":"### 4.4 Submission Visualization\nIn this section we will visualize the predicted flagellar coordinates in the tomographies","metadata":{}},{"id":"8ae8979f","cell_type":"code","source":"preferred_tomo_id = \"tomo_00e047\"\nif preferred_tomo_id in test_tomo_ids:\n    print(\"\\n--- Visualizing Test Tomography with Predictions ---\")\n    \n    tomo_id = preferred_tomo_id\n    print(f\"Using preferred tomography: {tomo_id}\")\n    print(f\"Processing visualization for: {tomo_id}\")\n\n    predicted_result = next((r for r in submission_results if r['tomo_id'] == tomo_id), None)\n    predicted_point = None\n    if predicted_result:\n        predicted_point = Point3D(\n            x=predicted_result['Motor axis 2'],\n            y=predicted_result['Motor axis 1'], \n            z=predicted_result['Motor axis 0']\n        )\n\n    actual_points = []\n    if tomo_id in train_df['tomo_id'].values:\n        print(f\"  Found {tomo_id} in training data - getting actual flagellar coordinates\")\n        actual_points = getFlagellarCoordinates(train_df, tomo_id)\n    else:\n        print(f\"  {tomo_id} not found in training data - no actual coordinates available\")\n\n    if predicted_point and predicted_point.x != -1 and predicted_point.y != -1 and predicted_point.z != -1:\n        print(f\"  Predicted point: {predicted_point}\")\n        \n        image_files = getImagesPath(base_dir, \"test\", tomo_id)\n        full_tomography = load_tomography(image_files)\n        \n        extraction_box = BoxShape(\n            x=EXPORT_SHAPE.x * 2,\n            y=EXPORT_SHAPE.y * 2,\n            z=EXPORT_SHAPE.z * 2,\n            center=predicted_point\n        )\n        \n        extraction_coords = extraction_box.get_box_coords()\n        \n        tomo_z, tomo_y, tomo_x = full_tomography.shape\n        x_start = max(0, int(extraction_coords.x_start))\n        x_end = min(tomo_x, int(extraction_coords.x_end))\n        y_start = max(0, int(extraction_coords.y_start))\n        y_end = min(tomo_y, int(extraction_coords.y_end))\n        z_start = max(0, int(extraction_coords.z_start))\n        z_end = min(tomo_z, int(extraction_coords.z_end))\n        \n        extracted_cube = full_tomography[z_start:z_end, y_start:y_end, x_start:x_end]\n        \n        resized_cube = resize_cuda_optimized(\n            extracted_cube, \n            EXPORT_SHAPE.toTuple(), \n            anti_aliasing=ANTIALIASING, \n            preserve_range=True, \n            device=DEVICE\n        ).astype(np.uint8)\n        \n        viz_points = []\n        \n        original_shape = extracted_cube.shape\n        scale_x = EXPORT_SHAPE.x / original_shape[2]\n        scale_y = EXPORT_SHAPE.y / original_shape[1]\n        scale_z = EXPORT_SHAPE.z / original_shape[0]\n        \n        pred_in_cube = Point3D(\n            x=(predicted_point.x - x_start) * scale_x,\n            y=(predicted_point.y - y_start) * scale_y,\n            z=(predicted_point.z - z_start) * scale_z\n        )\n        viz_points.append(pred_in_cube)\n        print(f\"  Predicted point in cube coordinates: {pred_in_cube}\")\n        \n        actual_points_info = []\n        if actual_points:\n            for j, actual_point in enumerate(actual_points):\n                actual_in_cube = Point3D(\n                    x=(actual_point.x - x_start) * scale_x,\n                    y=(actual_point.y - y_start) * scale_y,\n                    z=(actual_point.z - z_start) * scale_z\n                )\n                viz_points.append(actual_in_cube)\n                \n                is_in_bounds = (x_start <= actual_point.x <= x_end and \n                              y_start <= actual_point.y <= y_end and \n                              z_start <= actual_point.z <= z_end)\n                \n                actual_points_info.append({\n                    'original': actual_point,\n                    'in_cube': actual_in_cube,\n                    'in_bounds': is_in_bounds\n                })\n                \n                status = \"in bounds\" if is_in_bounds else \"outside bounds\"\n                print(f\"  Actual point {j+1}: {actual_point} ({status})\")\n                print(f\"    -> Cube coordinates: {actual_in_cube}\")\n        \n        title_parts = [f\"Test Tomography Cube: {tomo_id}\"]\n        title_parts.append(f\"Cube size: {EXPORT_SHAPE.toTuple()}\")\n        title_parts.append(f\"Extracted region: Z({z_start}-{z_end}), Y({y_start}-{y_end}), X({x_start}-{x_end})\")\n        title_parts.append(f\"Predicted: ({predicted_point.z:.0f}, {predicted_point.y:.0f}, {predicted_point.x:.0f})\")\n        \n        if actual_points_info:\n            in_bounds_count = sum(1 for info in actual_points_info if info['in_bounds'])\n            total_count = len(actual_points_info)\n            actual_coords = \", \".join([f\"({info['original'].z:.0f}, {info['original'].y:.0f}, {info['original'].x:.0f})\" \n                                     for info in actual_points_info])\n            title_parts.append(f\"Actual ({in_bounds_count}/{total_count} in bounds): {actual_coords}\")\n        else:\n            title_parts.append(\"Actual: Unknown (test data)\")\n        \n        title = \"\\n\".join(title_parts)\n        \n        display_tomography(resized_cube, title=title, flagellar_points=viz_points)\n        \n    else:\n        print(f\"  No valid prediction made for {tomo_id} - skipping visualization\")","metadata":{},"outputs":[],"execution_count":null},{"id":"9277e998","cell_type":"markdown","source":"### 4.5 Submission File Export\nIn this last cell we will export the submission file in the required format.","metadata":{}},{"id":"a6351fc1","cell_type":"code","source":"submission_file_path = \"submission.csv\"\n\nif submission_results:\n    \n    submission_df = pd.DataFrame(submission_results)\n    submission_df = submission_df[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]\n    submission_df.to_csv(submission_file_path, index=False)\n    print(f\"\\nSubmission file saved to: {os.path.abspath(submission_file_path)}\")\n\n    positive_predictions = len(submission_df[(submission_df['Motor axis 0'] != -1) & \n                                            (submission_df['Motor axis 1'] != -1) & \n                                            (submission_df['Motor axis 2'] != -1)])\n    print(f\"Positive predictions: {positive_predictions}/{len(submission_df)}\")\n    \n    print(\"\\n--- Submission DataFrame (first 10 rows) ---\")\n    print(submission_df.head(10))\n\n    if all_predictions_log:\n        predictions_df = pd.DataFrame(all_predictions_log)\n        predictions_file_path = \"detailed_predictions_log.csv\"\n        predictions_df.to_csv(predictions_file_path, index=False)\n        print(f\"Detailed predictions log saved to: {os.path.abspath(predictions_file_path)}\")\n          \nelse:\n    print(f\"\\nNo data in submission_results to save to {submission_file_path}. No predictions were made.\")\n\nprint(\"Submission generation completed.\")","metadata":{},"outputs":[],"execution_count":null}]}