{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from IPython.core.display import HTML\n\n\ndef apply_css_styles(\n    font_name: str = \"Lato\",\n    fallback_font: str = \"Verdana\",\n    css_content: str = None,\n    verbose: bool = False\n) -> HTML:\n    \"\"\"Applies custom CSS styles within a Jupyter notebook cell.\n\n    Args:\n        font_name (str, optional): \n            The primary font to use in the styles.\n        fallback_font (str, optional): \n            The fallback font to use if the primary font is unavailable.\n        css_content (str, optional): \n            Custom CSS content to use. \n            If None, default styles are used.\n        verbose (bool, optional): \n            Whether to print the generated CSS for debugging.\n\n    Returns:\n        IPython.core.display.HTML: \n            HTML object with the injected styles.\n    \"\"\"\n    try:\n        # Default CSS content if none is provided\n        default_css = '''\np, li, a, b, h1, h2, h3, h4, h5, h6, title, ul, strong, sup, sub, em, i, blockquote, label {\n    font-family: Verdana !important;\n}\n\nb, h1 {\n    font-weight: 900 !important;\n}\n\nh2, h3, h4 ul {\n    font-weight: 700 !important;\n}\n\n.fa, .far, .fas {\n    font-family: \"Font Awesome 5 Free\" !important;\n}\n'''\n\n        # Generate font import string dynamically based on the provided font name\n        font_import = (\n            f\"\\n@import url('https://fonts.googleapis.com/css2?family={font_name.replace(' ', '+')}:ital,wght@0,100;0,300;0,400;0,700;0,900;1,100;1,300;1,400;1,700;1,900&display=swap');\\n\"\n        )\n\n        # Use provided CSS content or fallback to default\n        css_to_use = css_content or default_css\n\n        # Replace fallback font in the CSS content\n        css_to_use = css_to_use.replace(\"Verdana\", font_name)\n\n        # Combine the font import and the CSS content into a single HTML style block\n        combined_styles = f\"<style>{font_import}{css_to_use}</style>\"\n\n        if verbose:\n            print(combined_styles)  # Print the CSS for debugging if verbose is True\n\n        return HTML(combined_styles)  # Return the generated styles as an HTML object\n\n    except Exception as e:\n        raise RuntimeError(f\"An error occurred while applying styles: {str(e)}\")\n\n# Apply styles (example usage)\napply_css_styles(verbose=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:49.797450Z","iopub.execute_input":"2025-03-14T13:13:49.797864Z","iopub.status.idle":"2025-03-14T13:13:49.809502Z","shell.execute_reply.started":"2025-03-14T13:13:49.797823Z","shell.execute_reply":"2025-03-14T13:13:49.808074Z"},"jupyter":{"source_hidden":true},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<div style=\"position: relative; width: 100%;\">\n  <img src=\"https://github.com/darien-schettler/asset-hosting/blob/main/fm_upscaled.png?raw=true\" width=\"100%\" style=\"padding: 0 0 !important; margin: 0 0 !important;\">\n  <div style=\"position: absolute; top: 5%; left: 3%; background-color: rgba(255, 255, 255, 0.85); padding: 20px; border-radius: 16px; box-shadow: 0 0 15px rgba(0,0,0,0.66);\">\n    <h1 style=\"text-align: center; font-size: 3.0vw !important; font-weight: 700; color: #6c0f1c; margin: 0; text-shadow: 1px 1px 3px rgba(0,0,0,0.25); letter-spacing: 2px; font-size: clamp(16px, 2.5vw, 44px);\">Part 1 - Data Processing</h1>\n  </div>\n</div>\n\n<br style=\"margin: 15px;\">\n\n<h2 style=\"text-align: center; font-size: 30px; font-style: normal; font-weight: 800; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">\n    <span style=\"text-decoration: underline;\">\n        <font color=#d27582>L</font>ET'S \n        <font color=#d27582>L</font>EARN \n        <font color=#d27582>T</font>OGETHER <font color=#d27582>!</font>!\n    </span><br><br><br style=\"margin: 15px;\">\n<span style=\"font-size: 22px; letter-spacing: 1px;\">\n    <font color=#d27582>U</font>NDERSTANDING    \n    <font color=#d27582>T</font>HROUGH\n    <font color=#d27582>E</font>XPLORATION\n</span><br style=\"margin: 15px;\"></h2>\n\n<p style=\"text-align: center; font-size: 15px; font-style: normal; font-weight: bold; text-decoration: None; text-transform: none; letter-spacing: 1px; color: black; background-color: #ffffff;\">CREATED BY: DARIEN SCHETTLER</p>\n\n<hr>\n\n<center><div class=\"alert alert-block alert-danger\" style=\"margin: 2em; line-height: 1.7em;\">\n    <b style=\"font-size: 18px;\">🛑 &nbsp; WARNING:</b><br><br><b>THIS IS A WORK IN PROGRESS<br><br><span style=\"color: #d27582\">📖 Many of the code cells will be compressed in the viewer to aid in readability 📖</span></b><br>\n</div></center>\n\n<center><div class=\"alert alert-block alert-warning\" style=\"margin: 2em; line-height: 1.7em;\">\n    <b style=\"font-size: 16px;\">👏 &nbsp; IF YOU FORK THIS OR FIND THIS HELPFUL &nbsp; 👏</b><br><br><b style=\"font-size: 22px; color: darkorange\">PLEASE UPVOTE!</b><br><br>This was a lot of work for me and while it may seem silly, it makes me feel appreciated when others like my work. 😅\n</div></center>\n\n<hr>","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: #6c0f1c; background-color: #ffffff;\">\n    CHANGELOG\n</h1>\n\n<ul>\n    <li>\n        <b>Version 1-8</b>\n        <ul>\n            <li>Initial Setup and Prep For Sharing</li>\n        </ul>\n    </li>\n</ul>\n\n<br>","metadata":{}},{"cell_type":"markdown","source":"<p id=\"toc\"></p>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: #6c0f1c; background-color: #ffffff;\">\n    TABLE OF CONTENTS\n</h1>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#introduction\" style=\"text-decoration: none; color: #d27582;\">1&nbsp;&nbsp;&nbsp;&nbsp;INTRODUCTION & JUSTIFICATION</a></h3>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#background_information\" style=\"text-decoration: none; color: #d27582;\">2&nbsp;&nbsp;&nbsp;&nbsp;BACKGROUND INFORMATION</a></h3>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#imports\" style=\"text-decoration: none; color: #d27582;\">3&nbsp;&nbsp;&nbsp;&nbsp;IMPORTS</a></h3>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#setup\" style=\"text-decoration: none; color: #d27582;\">4&nbsp;&nbsp;&nbsp;&nbsp;SETUP AND HELPER FUNCTIONS</a></h3>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#eda\" style=\"text-decoration: none; color: #d27582;\">5&nbsp;&nbsp;&nbsp;&nbsp;EXPLORATORY DATA ANALYSIS</a></h3>\n\n<hr>\n\n<h3 style=\"text-indent: 10vw; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; background-color: #ffffff;\"><a href=\"#preprocessing\" style=\"text-decoration: none; color: #d27582;\">6&nbsp;&nbsp;&nbsp;&nbsp;PROCESSING</a></h3>\n\n<hr>","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<a id=\"introduction\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"introduction\">1&nbsp;&nbsp;INTRODUCTION & JUSTIFICATION&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">1.1 <b>WHAT</b> IS THIS?</h3>\n<hr>\n\n<ul>\n    <li>This notebook explores the <b>Flagellar Motor Detection</b> challenge, in which participants must detect and locate <b>flagellar motors</b> in 3D bacterial tomograms generated via <b>cryo-electron tomography (cryo-ET)</b>.</li>\n    <li>We will walk through the core data, definitions, evaluation metrics, and potential approaches for building a solution.</li>\n    <li>We will also reference relevant biological terms (flagellar motors, proton gradients, rotor/stator assemblies, etc.) and imaging concepts (3D volumes, voxel spacing, noisy reconstructions).</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">1.2 <b>WHY</b> DOES THIS MATTER?</h3>\n<hr>\n\n<ul>\n    <li><b>Fundamental Biology</b>: Identifying the position of these molecular machines helps scientists study bacterial motility and energy usage.</li>\n    <li><b>Medical &amp; Drug Development</b>: Insights into motor function can guide new approaches for disrupting bacterial locomotion, which may help combat infections and reduce antibiotic resistance.</li>\n    <li><b>Accelerating Cryo-ET Workflows</b>: Manually labeling tomograms is time-consuming. Automated detection can speed up data analysis, enabling larger-scale experiments.</li>\n    <li><b>Machine Learning Innovation</b>: This is a challenging 3D object detection problem (low SNR, large volumes, variable orientations), so novel solutions can push forward the state of the art in volumetric image analysis.</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">1.3 <b>WHO</b> IS THIS FOR?</h3>\n<hr>\n\n<ul>\n    <li>These notebooks are always primarily for <b>me</b>... however, there are other individuals or groups who could benefit:</li>\n        <ul>\n            <li>Biologists and Biophysicists eager to automate the analysis of cryo-ET data.</li>\n            <li>Data Scientists and Computer Vision practitioners interested in new challenges in 3D detection and segmentation.</li>\n            <li>Anyone curious about bridging the gap between AI and advanced microscopy techniques.</li>\n        </ul>\n    </li>\n    <li>Whether you're from a biology background learning ML or an ML enthusiast picking up domain knowledge, this notebook aims to help you make meaningful progress.</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">1.4 <b>HOW</b> WILL THIS WORK?</h3>\n<hr>\n\n<p>We'll begin with a brief overview of the data and a review of key biological and cryo-ET concepts. Next, we’ll walk through four stages (1 per notebook):</p>\n<ul>\n    <li><a href=\"https://www.kaggle.com/code/dschettler8845/byu-1-4-data-processing-let-s-learn-together\">[THIS NOTEBOOK] </a> Exploratory Data Analysis followed by Data Preprocessing for YOLO.</li>\n    <li><a href=\"https://www.kaggle.com/code/dschettler8845/byu-2-4-data-processing-let-s-learn-together\">[NOTEBOOK 2]</a> <b>TBD</b>.</li>\n    <li><a href=\"https://www.kaggle.com/code/dschettler8845/byu-3-4-data-processing-let-s-learn-together\">[NOTEBOOK 3]</a> <b>TBD</b>.</li>\n    <li><a href=\"https://www.kaggle.com/code/dschettler8845/byu-4-4-data-processing-let-s-learn-together\">[NOTEBOOK 4]</a> <b>TBD</b>.</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<a id=\"background_information\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"background_information\">2&nbsp;&nbsp;BACKGROUND INFORMATION&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n\nThe <b>flagellar motor</b> is a biological nanomachine that propels many bacteria through fluids. It is embedded across the bacterial cell membranes and is powered by ionic gradients (usually protons). The motor rotates a helical flagellum, enabling processes like <b>chemotaxis</b> and pathogenesis. In <b>cryo-ET</b>, we can capture these motors in near-native states, providing detailed views of their structures but also introducing challenges in data analysis due to noisy, high-dimensional images.\n\nThis competition leverages <b>3D reconstructions</b> (tomograms) of bacteria, each possibly containing <b>0, 1, or multiple</b> motors. \n    \nOur task: <b>automatically detect if a motor is present, and if so, localize it by returning the x, y, z coordinates</b>.\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">2.1 <b>COMPETITION OVERVIEW</b></h3>\n<hr>\n\n<br>\n\n<b style=\"text-decoration: underline; font-size: 15px; text-transform: uppercase; letter-spacing: 2px; font-weight: 900;\">PRIMARY TASK DESCRIPTION</b>\n<br>\nDevelop a 3D image processing algorithm to detect whether a flagellar motor exists in a tomogram, and if so, predict its center coordinates (Motor axis 0, Motor axis 1, Motor axis 2). \n<br><br>\n<b style=\"text-decoration: underline; font-size: 15px; text-transform: uppercase; letter-spacing: 2px; font-weight: 900;\">KEY DETAILS</b>\n<br>\n<ul>\n    <li><b>Input Data:</b> ~800 3D tomograms split into 2D slices (JPEG format). Each <code>tomo_id</code> has its own subdirectory.</li>\n    <li><b>Training Labels:</b> Provided in <code>train_labels.csv</code> with the annotated coordinates (if any) for each motor, plus shape info for the tomogram.</li>\n    <li><b>Test Data:</b> ~900 tomograms, each containing either 0 or 1 motor.</li>\n    <li><b>Output:</b> A <code>submission.csv</code> with the predicted coordinates for each <code>tomo_id</code>. Use <code>-1,-1,-1</code> if no motor is predicted.</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">2.2 <b>EVALUATION METRICS</b></h3>\n<hr>\n\n<p><b>Submissions</b> are evaluated via a combination of:</p>\n<ul>\n    <li><b>Euclidean Distance</b> - Measures how close your predicted location is to the true motor location. If the distance is &le; 1000 Å, you earn a <i>true positive</i> for that tomogram.</li>\n    <li><b>F<sub>2</sub>-score</b> (F<sub>β</sub> with β=2) - Balances <b>precision</b> and <b>recall</b>, placing greater emphasis on recall. In other words, missing an actual motor (a false negative) is penalized more heavily than predicting an extra one (a false positive).</li>\n</ul>\n\nThe final score rewards both correctly identifying motors and pinpointing their positions with high accuracy. \n\nIf you predict no motor exists, set <code>Motor axis 0, Motor axis 1, Motor axis 2 = -1</code>.\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">2.3 <b>DATASET INFORMATION</b></h3>\n<hr>\n\n<p><b style=\"text-decoration: underline; font-size: 15px; text-transform: uppercase; letter-spacing: 2px; font-weight: 900;\">HIGH LEVEL DATA SUMMARY</b></p>\n<ul>\n    <li>3D volumes provided as <i>stacks of 2D slices</i> (JPEG images). \n    <li>Training labels in <code>train_labels.csv</code> – each row represents a motor or a “no motor” label with <code>-1</code> coords.</li>\n    <li>817 training tomograms (some with multiple motors), ~900 test tomograms (each with at most one motor).</li>\n    <li>Large data size (~73.87 GB). <b>Voxel spacing</b> in Angstroms is crucial to interpret real-world distances.</li>\n</ul>\n\n<p><b style=\"text-decoration: underline; font-size: 15px; text-transform: uppercase; letter-spacing: 2px; font-weight: 900;\">DATA FILE DESCRIPTIONS</b></p>\n<ul>\n    <li><code>train/</code>: Subdirectories of 2D slices (JPEG) for each tomogram <code>tomo_id</code>. Label info in <code>train_labels.csv</code>.</li>\n    <li><code>train_labels.csv</code>: Coordinates (z, y, x) for each motor or <code>-1, -1, -1</code> if none. Includes tomogram shape and voxel spacing.</li>\n    <li><code>test/</code>: Subdirectories for test tomograms (also 2D slices). In the hidden test set, each <code>tomo_id</code> has 0 or 1 motor.</li>\n    <li><code>sample_submission.csv</code>: Format template. Provide coordinates for each <code>tomo_id</code>, or <code>-1, -1, -1</code> if no motor is predicted.</li>\n</ul>\n\n<p><b style=\"text-decoration: underline; font-size: 15px; text-transform: uppercase; letter-spacing: 2px; font-weight: 900;\">DATASET CHALLENGES</b></p>\n<ul>\n    <li>Extremely <b>noisy 3D data</b>; requires robust feature extraction. </li>\n    <li><b>Variable motor orientation</b> and potential partial views near tomogram edges.</li>\n    <li><b>Limited &amp; labor-intensive annotations</b> – each motor location must be precisely labeled, increasing data scarcity.</li>\n</ul>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">2.4 <b>MY COMPETITION DIARY/WEBSITE WITH LOTS OF ADDITIONAL DETAIL &amp; TIPS</b></h3>\n<hr>\n\n<br>\n\nPlease see my notion if you'd like to learn more:\n\n---\n\n<a href=\"https://striped-toothpaste-66b.notion.site/BYU-Locating-Bacterial-Flagellar-Motors-2025-1b129fe53ebd80fd92d0ee8b957bc3a0?pvs=4\"><b style=\"font-size: 24px;\">MT NOTION THAT WILL CONTINUALLY BE UPDATED WITH MY LEARNINGS</b></a>\n\n---\n\n<br>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<a id=\"imports\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"imports\">3&nbsp;&nbsp;IMPORTS&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n","metadata":{}},{"cell_type":"code","source":"print(\"\\n... PIP INSTALLS STARTING ...\\n\")\nprint(\"\\n... PIP INSTALLS COMPLETE ...\\n\")\n\nprint(\"\\n... IMPORTS STARTING ...\\n\")\nprint(\"\\n\\tVERSION INFORMATION\")\nimport pandas as pd; pd.options.mode.chained_assignment = None; pd.set_option('display.max_columns', None);\nimport numpy as np; print(f\"\\t\\t– NUMPY VERSION: {np.__version__}\");\nimport sklearn; print(f\"\\t\\t– SKLEARN VERSION: {sklearn.__version__}\");\n\n# Built-In Imports (mostly don't worry about these)\nfrom typing import Iterable, Any, Callable, Generator\nfrom kaggle_datasets import KaggleDatasets\nfrom dataclasses import dataclass\nfrom collections import Counter\nfrom datetime import datetime\nfrom zipfile import ZipFile\nfrom glob import glob\nimport subprocess\nimport warnings\nimport requests\nimport textwrap\nimport hashlib\nimport imageio\nimport IPython\nimport urllib\nimport zipfile\nimport pickle\nimport random\nimport shutil\nimport string\nimport yaml\nimport json\nimport copy\nimport math\nimport time\nimport gzip\nimport ast\nimport sys\nimport io\nimport gc\nimport re\nimport os\n\n# Visualization Imports (overkill)\nfrom IPython.core.display import HTML, Markdown\nfrom matplotlib.patches import Rectangle\nimport matplotlib.colors as mcolors\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm; tqdm.pandas();\nfrom mpl_toolkits.mplot3d import Axes3D\nimport plotly.express as px\nimport plotly.graph_objects as go\nimport seaborn as sns\nfrom PIL import Image, ImageEnhance; Image.MAX_IMAGE_PIXELS = 5_000_000_000;\nimport matplotlib; print(f\"\\t\\t– MATPLOTLIB VERSION: {matplotlib.__version__}\");\nimport plotly\nimport PIL\n\n# Rich\nimport rich\nfrom rich import pretty; pretty.install()\nfrom rich.markdown import Markdown\nfrom rich import print as rprint\nfrom rich.console import Console\nfrom rich.style import Style\nfrom rich.live import Live\nfrom rich.text import Text\nfrom rich import inspect\n\ndef seed_it_all(seed=7):\n    \"\"\" Attempt to be Reproducible \"\"\"\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    # tf.random.set_seed(seed)\n    \nseed_it_all()\n\nwarnings.filterwarnings('ignore', category=FutureWarning, message='use_inf_as_na option is deprecated')\nprint(\"\\n\\n... IMPORTS COMPLETE ...\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:49.810857Z","iopub.execute_input":"2025-03-14T13:13:49.811162Z","iopub.status.idle":"2025-03-14T13:13:52.729974Z","shell.execute_reply.started":"2025-03-14T13:13:49.811133Z","shell.execute_reply":"2025-03-14T13:13:52.728417Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<a id=\"setup\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"setup\">4&nbsp;&nbsp;SETUP AND HELPER FUNCTIONS&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">4.0 FUNCTIONS FROM <b>OTHER KAGGLERS</b> ♥️</h3>\n<hr><br>","metadata":{}},{"cell_type":"code","source":"\"\"\"The code below is heavily inspired by the notebook: https://www.kaggle.com/code/andrewjdarley/parse-data\n\n    (1) I first understood Andrew's code, then I rewrote it to be more aligned with my own style.\n    (2) Next, I updated it to incorporate any changes I think are appropriate.\n    (3) Last, in the cell in section 6, I wrap this into a single unified function to run the prep end to end.\n\"\"\"\n\n\ndef create_dataset_directories(\n    yolo_images_train: str, \n    yolo_images_val: str, \n    yolo_labels_train: str, \n    yolo_labels_val: str\n) -> None:\n    \"\"\"\n    Create all necessary directories for the YOLO dataset.\n    \n    Args:\n        yolo_images_train (str): Path to directory for training images\n        yolo_images_val (str): Path to directory for validation images\n        yolo_labels_train (str): Path to directory for training labels\n        yolo_labels_val (str): Path to directory for validation labels\n    \"\"\"\n    for dir_path in [yolo_images_train, yolo_images_val, yolo_labels_train, yolo_labels_val]:\n        os.makedirs(dir_path, exist_ok=True)\n        print(f\"CREATED DIRECTORY\\n\\t--> {dir_path}\")\n\n\ndef normalize_slice(slice_data: np.ndarray) -> np.ndarray:\n    \"\"\"Normalize a tomographic slice for better visualization and learning.\n    \n    Uses percentile-based normalization to enhance contrast while \n    preserving important features.\n    \n    Args:\n        slice_data (np.ndarray): \n            Raw numpy array of the tomographic slice\n        \n    Returns:\n        np.ndarray:\n            Normalized slice data as uint8 numpy array (0-255 range)\n    \"\"\"\n    # (1) Calculate 2nd and 98th percentiles for robust normalization\n    p2 = np.percentile(slice_data, 2)\n    p98 = np.percentile(slice_data, 98)\n    \n    # (2) Clip the data to the percentile range to reduce outlier influence\n    clipped_data = np.clip(slice_data, p2, p98)\n    \n    # (3) Normalize to [0, 255] range for standard image representation\n    normalized = 255 * (clipped_data - p2) / (p98 - p2)\n    \n    # (4) Convert to 8-bit unsigned integer format for image saving\n    return np.uint8(normalized)\n    \n\ndef validate_labels_file(labels_path: str) -> pd.DataFrame:\n    \"\"\"Load and validate the labels CSV file.\n    \n    Args:\n        labels_path (str): Path to the labels CSV file\n        \n    Returns:\n        pd.DataFrame:\n            The processed label information.\n        \n    Raises:\n        FileNotFoundError: If labels file doesn't exist\n        ValueError: If labels file format is invalid\n    \"\"\"\n    # (0) Define the required columns\n    _required_columns = [\n        'tomo_id', \n        'Motor axis 0', \n        'Motor axis 1', \n        'Motor axis 2',\n        'Array shape (axis 0)', \n        'Number of motors'\n    ]\n    \n    # (1) Check if labels file exists\n    if not os.path.exists(labels_path):\n        raise FileNotFoundError(f\"Labels file not found: {labels_path}\")\n    \n    # (2) Load the labels CSV\n    try:\n        labels_df = pd.read_csv(labels_path)\n    except Exception as e:\n        raise ValueError(f\"Error reading labels file: {e}\")\n    \n    # (3) Validate required columns exist\n    missing_columns = [col for col in _required_columns if col not in labels_df.columns]\n    if missing_columns:\n        raise ValueError(f\"Labels file missing required columns: {', '.join(missing_columns)}\")\n        \n    # (4) Return the validated DataFrame\n    return labels_df\n\n\ndef split_tomograms(\n    labels_df: pd.DataFrame, \n    train_split: float = 0.8,\n    random_seed: int = 42\n) -> tuple[np.ndarray, np.ndarray]:\n    \"\"\"Split tomograms into training and validation sets.\n    \n    Performs split at the tomogram level to ensure all slices from a single\n    tomogram are in the same set (train or validation).\n    \n    Args:\n        labels_df (pd.DataFrame): \n            DataFrame containing the tomogram labels\n        train_split (float, optional): \n            Fraction of data to use for training (0.0-1.0)\n        random_seed (int, optional): \n            Random seed for reproducibility\n        \n    Returns:\n        tuple[np.ndarray, np.ndarray]:\n            The tomograms split into their respective train and validation distributions.\n    \"\"\"\n    # (1) Set random seed for reproducibility\n    np.random.seed(random_seed)\n    \n    # (2) Find tomograms that have motors (tomograms of interest)\n    tomo_df = labels_df[labels_df['Number of motors'] > 0].copy()\n    unique_tomos = tomo_df['tomo_id'].unique()\n    \n    print(f\"\\nFOUND {len(unique_tomos)} UNIQUE TOMOGRAMS WITH 1 OR MORE MOTORS\\n\")\n    \n    # (3) Shuffle tomograms for random split\n    np.random.shuffle(unique_tomos)\n    \n    # (4) Calculate split index based on train_split ratio\n    split_idx = int(len(unique_tomos) * train_split)\n    \n    # (5) Create train and validation sets\n    train_tomos = unique_tomos[:split_idx]\n    val_tomos = unique_tomos[split_idx:]\n    \n    print(f\"\\nSPLIT DISTRIBUTION:\\n\\tTRAIN: {len(train_tomos)} TOMOGRAMS\\n\\tVALIDATION: {len(val_tomos)} TOMOGRAMS.\")\n    \n    return train_tomos, val_tomos\n\ndef process_tomogram_set(\n    labels_df: pd.DataFrame,\n    tomogram_ids: np.ndarray, \n    train_dir: str,\n    images_dir: str, \n    labels_dir: str, \n    set_name: str,\n    trust: int = 4,\n    box_size: int = 24\n) -> tuple[int, int]:\n    \"\"\"Process a set of tomograms, extracting slices and creating annotations.\n    \n    Args:\n        labels_df (pd.DataFrame): DataFrame containing the tomogram labels\n        tomogram_ids (np.ndarray): Array of tomogram IDs to process\n        train_dir (str): Directory containing the raw tomogram data\n        images_dir (str): Directory to save processed images\n        labels_dir (str): Directory to save annotation labels\n        set_name (str): Name of the dataset (e.g., \"training\" or \"validation\")\n        trust (int, optional): Number of slices above and below center slice to include\n        box_size (int, optional): Size of bounding box in pixels for annotations\n        \n    Returns:\n        tuple[int, int]:\n            The count of slices and the count of motors.\n    \"\"\"\n    # (1) Extract motor information for the specified tomograms\n    motor_info = []\n    for tomo_id in tomogram_ids:\n        # Get all motors for this tomogram\n        tomo_motors = labels_df[labels_df['tomo_id'] == tomo_id]\n        for _, motor in tomo_motors.iterrows():\n            if pd.isna(motor['Motor axis 0']):\n                continue\n            motor_info.append(\n                (tomo_id, \n                 int(motor['Motor axis 0']), \n                 int(motor['Motor axis 1']), \n                 int(motor['Motor axis 2']),\n                 int(motor['Array shape (axis 0)']))\n            )\n    \n    # (2) Output processing information\n    print(f\"\\nPROCESSING APPROXIMATELY {len(motor_info) * (2 * trust + 1)} SLICES FOR '{set_name}'\\n\")\n    \n    # (3) Process each motor\n    processed_slices = 0\n    \n    # (4) Process all motors across all tomograms in the set\n    for tomo_id, z_center, y_center, x_center, z_max in tqdm(motor_info, desc=f\"PROCESSING {set_name} MOTORS\"):\n        # (4.1) Calculate range of slices to include based on trust parameter\n        z_min = max(0, z_center - trust)\n        z_max = min(z_max - 1, z_center + trust)\n        \n        # (4.2) Process each slice in the defined range\n        for z in range(z_min, z_max + 1):\n            # Create slice filename\n            slice_filename = f\"slice_{z:04d}.jpg\"\n            \n            # Source path for the slice\n            src_path = os.path.join(train_dir, tomo_id, slice_filename)\n            \n            # (4.3) Skip if source file doesn't exist\n            if not os.path.exists(src_path):\n                print(f\"Warning: {src_path} does not exist, skipping.\")\n                continue\n            \n            # (4.4) Load and normalize the slice\n            try:\n                img = Image.open(src_path)\n                img_array = np.array(img)\n            except Exception as e:\n                print(f\"Error loading image {src_path}: {e}\")\n                continue\n            \n            # (4.5) Normalize the image\n            try:\n                normalized_img = normalize_slice(img_array)\n            except Exception as e:\n                print(f\"Error normalizing image {src_path}: {e}\")\n                continue\n            \n            # (4.6) Create destination filename with unique identifier\n            dest_filename = f\"{tomo_id}_z{z:04d}_y{y_center:04d}_x{x_center:04d}.jpg\"\n            dest_path = os.path.join(images_dir, dest_filename)\n            \n            # (4.7) Save the normalized image\n            try:\n                Image.fromarray(normalized_img).save(dest_path)\n            except Exception as e:\n                print(f\"Error saving image {dest_path}: {e}\")\n                continue\n            \n            # (4.8) Get image dimensions for normalization\n            img_width, img_height = img.size\n            \n            # (4.9) Create YOLO format label (see below)\n            #    - <class> <x_center> <y_center> <width> <height>\n            #    - Values are normalized to [0, 1]\n            x_center_norm = x_center / img_width\n            y_center_norm = y_center / img_height\n            box_width_norm = box_size / img_width\n            box_height_norm = box_size / img_height\n            \n            # (4.10) Write label file\n            label_path = os.path.join(labels_dir, dest_filename.replace('.jpg', '.txt'))\n            try:\n                with open(label_path, 'w') as f:\n                    f.write(f\"0 {x_center_norm} {y_center_norm} {box_width_norm} {box_height_norm}\\n\")\n            except Exception as e:\n                print(f\"Error writing label {label_path}: {e}\")\n                continue\n            \n            # (4.11) Increment slice counter\n            processed_slices += 1\n    \n    # (5) Return statistics\n    return processed_slices, len(motor_info)\n\n\ndef create_yaml_config(yolo_dataset_dir: str) -> str:\n    \"\"\"Create YAML configuration file for YOLO training.\n    \n    Args:\n        yolo_dataset_dir (str): Base directory for the YOLO dataset\n        \n    Returns:\n        str: The path to the created YAML file\n    \"\"\"\n    # (1) Define YAML content with dataset paths and class names\n    yaml_content = {\n        'path': yolo_dataset_dir,\n        'train': 'images/train',\n        'val': 'images/val',\n        'names': {0: 'motor'}\n    }\n    \n    # (2) Define output path\n    yaml_path = os.path.join(yolo_dataset_dir, 'dataset.yaml')\n    \n    # (3) Write YAML file\n    try:\n        with open(yaml_path, 'w') as f:\n            yaml.dump(yaml_content, f, default_flow_style=False)\n    except Exception as e:\n        print(f\"Warning: Failed to write YAML config: {e}\")\n        \n    # (4) Return path to the created file\n    return yaml_path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:52.731938Z","iopub.execute_input":"2025-03-14T13:13:52.732652Z","iopub.status.idle":"2025-03-14T13:13:52.755354Z","shell.execute_reply.started":"2025-03-14T13:13:52.732611Z","shell.execute_reply":"2025-03-14T13:13:52.753982Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">4.1 GENERIC HELPER FUNCTIONS</h3>\n<hr><br>\n\nThese are some functions I carry around with me that I find commonly helpful.\n\n**There are also a few dataset loading functions included here just to ensure they fall before the setup cell**\n\n<br>","metadata":{}},{"cell_type":"code","source":"def flatten_l_o_l(nested_iterable: Iterable[Iterable[Any]]) -> list[Any]:\n    \"\"\"Flatten a list of lists (or any nested iterable) into a single list.\n    \n    Transforms a nested structure like [[1, 2], [3, 4]] into [1, 2, 3, 4].\n    \n    Args:\n        nested_iterable (Iterable[Iterable[Any]]): \n            An iterable containing other iterables to be flattened.\n            Examples: List of lists, tuple of sets, etc.\n    \n    Returns:\n        list[T]: A flattened list containing all items from the input nested structure.\n    \n    Examples:\n        >>> flatten_l_o_l([[1, 2], [3, 4]])\n        [1, 2, 3, 4]\n        >>> flatten_l_o_l([(5, 6), [7, 8]])\n        [5, 6, 7, 8]\n    \"\"\"\n    # (1) Use list comprehension with nested loops to flatten the structure\n    return [item for sublist in nested_iterable for item in sublist]\n\n\ndef print_ln(\n    symbol: str = \"-\", \n    line_len: int = 110, \n    newline_before: bool = False, \n    newline_after: bool = False\n) -> None:\n    \"\"\"Print a horizontal line of a specified length and symbol.\n    \n    Creates a visual separator in console output for improved readability.\n    \n    Args:\n        symbol: The character(s) to use for the horizontal line.\n            Single character strings work best (e.g., \"-\", \"=\", \"*\").\n        line_len: The length of the horizontal line in characters.\n            Default is 110 characters.\n        newline_before: Whether to print a newline character before the line.\n            Used to create spacing before the separator.\n        newline_after: Whether to print a newline character after the line.\n            Used to create spacing after the separator.\n            \n    Returns:\n        None: This function prints to stdout but doesn't return any value.\n    \n    Examples:\n        >>> print_ln()  # Prints \"----------...\" (110 dashes)\n        >>> print_ln(\"=\", 50, True, True)  # Prints a newline, then 50 \"=\" characters, then a newline\n    \"\"\"\n    # (1) Print a newline before the line if requested\n    if newline_before:\n        print()\n    \n    # (2) Print the line using string multiplication\n    print(symbol * line_len)\n    \n    # (3) Print a newline after the line if requested\n    if newline_after:\n        print()\n        \n        \ndef display_hr(\n    newline_before: bool = False, \n    newline_after: bool = False\n) -> None:\n    \"\"\"Display an HTML horizontal rule (<hr>) in notebook environments.\n    \n    Creates a visual separator in Jupyter/IPython notebook output.\n    \n    Args:\n        newline_before: Whether to print a newline character before the horizontal rule.\n            Used to create spacing before the separator.\n        newline_after: Whether to print a newline character after the horizontal rule.\n            Used to create spacing after the separator.\n            \n    Returns:\n        None: This function displays HTML content but doesn't return any value.\n    \n    Notes:\n        - This function is designed for use in Jupyter notebook or IPython environments.\n        - It will not render correctly in standard console environments.\n    \n    Examples:\n        >>> display_hr()  # Displays an HTML horizontal rule\n        >>> display_hr(True, True)  # Displays a newline, then an HTML horizontal rule, then a newline\n    \"\"\"\n    # (1) Print a newline before the HTML horizontal rule if requested\n    if newline_before:\n        print()\n    \n    # (2) Display the HTML horizontal rule\n    display(HTML(\"<hr>\"))\n    \n    # (3) Print a newline after the HTML horizontal rule if requested\n    if newline_after:\n        print()\n\n\ndef wrap_text(text: str, width: int = 88) -> str:\n    \"\"\"Wrap text to a specified width.\n    \n    Formats a long string by inserting line breaks to ensure no line exceeds\n    the specified width. Useful for formatting paragraphs for display in\n    fixed-width contexts.\n    \n    Args:\n        text: The text string to wrap.\n            Can be a single line or multiple lines.\n        width: The maximum width of a line in characters.\n            Default is 88 characters, which matches Black formatter's default.\n\n    Returns:\n        str: The wrapped text with added line breaks.\n    \n    Examples:\n        >>> long_text = \"This is a very long string that needs to be wrapped to multiple lines.\"\n        >>> wrap_text(long_text, 20)\n        'This is a very long\\\\nstring that needs to\\\\nbe wrapped to\\\\nmultiple lines.'\n    \"\"\"\n    # (1) Use textwrap.fill to wrap the text to the specified width\n    return textwrap.fill(text, width)\n\n\ndef wrap_text_by_paragraphs(text: str, width: int = 88) -> str:\n    \"\"\"Wrap text by paragraphs to a specified width while preserving paragraph structure.\n    \n    Similar to wrap_text(), but maintains paragraph separation by preserving\n    blank lines between paragraphs.\n    \n    Args:\n        text: The text string containing multiple paragraphs to wrap.\n            Paragraphs should be separated by newline characters.\n        width: The maximum width of a line in characters.\n            Default is 88 characters, which matches Black formatter's default.\n\n    Returns:\n        str: The wrapped text with preserved paragraph separation.\n    \n    Examples:\n        >>> paragraphs = \"First paragraph.\\\\n\\\\nSecond paragraph that is longer.\"\n        >>> wrap_text_by_paragraphs(paragraphs, 20)\n        'First paragraph.\\\\n\\\\nSecond paragraph\\\\nthat is longer.'\n    \"\"\"\n    # (1) Split the text into paragraphs using newline characters\n    paragraphs = text.split('\\n')\n    \n    # (2) Wrap each paragraph individually\n    wrapped_paragraphs = [textwrap.fill(paragraph, width) for paragraph in paragraphs]\n    \n    # (3) Join the wrapped paragraphs with double newlines to preserve paragraph structure\n    return '\\n\\n'.join(wrapped_paragraphs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:52.757119Z","iopub.execute_input":"2025-03-14T13:13:52.757466Z","iopub.status.idle":"2025-03-14T13:13:52.785579Z","shell.execute_reply.started":"2025-03-14T13:13:52.757425Z","shell.execute_reply":"2025-03-14T13:13:52.784260Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">4.2 <b>LOAD</b> THE DATASET(S)</h3>\n<hr><br>\n\nWe also define path information and other constants that are helpful in establishing early.\n\n<br>","metadata":{}},{"cell_type":"code","source":"# ROOT PATHS (define these in your notebook)\nWORKING_DIR = \"/kaggle/working\"\nINPUT_DIR = \"/kaggle/input\"\nCOMPETITION_DIR = os.path.join(INPUT_DIR, \"byu-locating-bacterial-flagellar-motors-2025\")\n\n# COMPETITION DATA PATHS\nTRAIN_DIR = os.path.join(COMPETITION_DIR, \"train\")\nTRAIN_LABELS_PATH = os.path.join(COMPETITION_DIR, \"train_labels.csv\")\nTEST_DIR = os.path.join(COMPETITION_DIR, \"test\")\n\n# OUTPUT PATHS\nYOLO_DATASET_DIR = os.path.join(WORKING_DIR, \"yolo_dataset\")\nYOLO_IMAGES_TRAIN = os.path.join(YOLO_DATASET_DIR, \"images\", \"train\")\nYOLO_IMAGES_VAL = os.path.join(YOLO_DATASET_DIR, \"images\", \"val\")\nYOLO_LABELS_TRAIN = os.path.join(YOLO_DATASET_DIR, \"labels\", \"train\")\nYOLO_LABELS_VAL = os.path.join(YOLO_DATASET_DIR, \"labels\", \"val\")\n\n# DATASET PROCESSING HYPERPARAMETERS\nTRUST = 4          # Number of slices above and below center slice\nBOX_SIZE = 24      # Bounding box size for annotations\nTRAIN_SPLIT = 0.8  # 80% for training, 20% for validation\n\n# LOAD THE DATASET\nlabels_df = validate_labels_file(TRAIN_LABELS_PATH)\n\nrich.print(\"\\n\\n[bold red]LABELS DATAFRAME[/bold red]\")\nlabels_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:52.786914Z","iopub.execute_input":"2025-03-14T13:13:52.787449Z","iopub.status.idle":"2025-03-14T13:13:52.874762Z","shell.execute_reply.started":"2025-03-14T13:13:52.787407Z","shell.execute_reply":"2025-03-14T13:13:52.873594Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<a id=\"eda\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"eda\">5&nbsp;&nbsp;EXPLORATORY DATA ANALYSIS&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n\n**Key Insights:**\n  - **No missing values**\n  - Total Motors: **737**\n  - Total Unique TomoGrams: **648**\n  - Average Motors per Tomogram: **0.70**\n  - Tomograms with Multiple Motors: **49**\n  - Average Dimensions (Z×Y×X): **422.7 × 950.2 × 954.8**\n  - Average Voxel Spacing: **15.34 Å**\n\n<br>","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">5.1 <b><code>BASIC</code></b> EXPLORATION</h3>\n<hr><br>\n\n<br>","metadata":{}},{"cell_type":"code","source":"labels_df.info()\nlabels_df.describe().T","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:52.875824Z","iopub.execute_input":"2025-03-14T13:13:52.876200Z","iopub.status.idle":"2025-03-14T13:13:52.936174Z","shell.execute_reply.started":"2025-03-14T13:13:52.876174Z","shell.execute_reply":"2025-03-14T13:13:52.935038Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">5.2 <b><code>TOMO_ID</code></b> EXPLORATION</h3>\n<hr><br>\n\n`tomo_id` is a unique identifier of the tomogram. Some tomograms in the train set have multiple motors.\n\n\n<br>","metadata":{}},{"cell_type":"code","source":"def visualize_motor_counts(\n    df: pd.DataFrame,\n    color_sequence: list[str] | None = None,\n    height: int = 500,\n    width: int = 800\n) -> px.pie:\n    \"\"\"\n    Creates a pie chart visualization showing the distribution of tomograms \n    by their motor count.\n    \n    Args:\n        df (pd.Dataframe): The tomogram dataset.\n        color_sequence (list[str], optional): Colors for the pie chart (default: Plotly default)\n        height (int, optional): Height of the figure in pixels.\n        width (int, optional): Width of the figure in pixels.\n        \n    Returns:\n        A Plotly pie chart figure showing distribution of tomograms by motor count\n        \n    Example:\n        >>> fig = visualize_motor_counts(labels_df)\n        >>> fig.show(renderer=\"iframe\")\n    \"\"\"\n    # (1) Group by tomo_id and get the number of motors for each unique tomogram\n    # We only need one row per tomogram since 'Number of motors' is the same for all rows of the same tomogram\n    motors_per_tomo = df.drop_duplicates(subset=['tomo_id'])[['tomo_id', 'Number of motors']]\n    \n    # (2) Count tomograms by their motor count\n    motor_count_distribution = motors_per_tomo['Number of motors'].value_counts().reset_index()\n    motor_count_distribution.columns = ['motor_count', 'num_tomograms']\n    \n    # (3) Sort by motor count for better interpretation\n    motor_count_distribution = motor_count_distribution.sort_values('motor_count')\n    \n    # (4) Create the pie chart\n    fig = px.pie(\n        motor_count_distribution, values='num_tomograms', names='motor_count',  # Information to plot\n        title='<b>Distribution of Tomograms by Motor Count</b>',                # Title\n        color_discrete_sequence=color_sequence,                                 # Colour Sequence\n        height=height, width=width, hole=0.3,                                   # Sizing and Whatnot\n        category_orders={\"motor_count\": sorted(motor_count_distribution['motor_count'].tolist())},\n    )\n    \n    # (5) Improve layout for better readability\n    fig.update_layout(\n        margin=dict(l=20, r=120, t=100, b=20),  # Increased right margin for legend\n        legend_title='<b>Number of Motors</b>',  # Bold legend title\n    )\n    \n    # (6) Add percentage and count to hover information and adjust text position\n    fig.update_traces(\n        textinfo='percent+label',\n        texttemplate='<b>%{label}</b><br>%{percent}',  # Bold labels\n        textposition='inside',                         # This avoids the callout interfering with the title.\n        hovertemplate='<b>Number of Motors: %{label}</b><br>Number of Tomograms: %{value}<br>Percentage: %{percent}'\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:13:52.938394Z","iopub.execute_input":"2025-03-14T13:13:52.938838Z","iopub.status.idle":"2025-03-14T13:13:54.860504Z","shell.execute_reply.started":"2025-03-14T13:13:52.938799Z","shell.execute_reply":"2025-03-14T13:13:54.859448Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"motor_counts_fig = visualize_motor_counts(labels_df)\nmotor_counts_fig.show(renderer='iframe')","metadata":{"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_tomo_distribution(\n    df: pd.DataFrame, \n    n_top: int = 10, \n    color: str = '#d27582',\n    height: int = 700,\n    width: int = 900\n) -> px.bar:\n    \"\"\"Visualizes the distribution of tomogram IDs in the dataset.\n    \n    This function counts the frequency of each tomogram ID and plots the top N most \n    frequently occurring tomograms, which correlates with those containing multiple motors.\n    \n    Args:\n        df (pd.DataFrame): DataFrame containing the tomogram dataset with at least 'tomo_id' column\n        n_top (int, optional): Number of top tomograms to display (default: 15)\n        color (str, optional): Color for the bar chart (default: #d27582)\n        height (int, optional): Height of the figure in pixels (default: 500)\n        width (int, optional): Width of the figure in pixels (default: 800)\n        \n    Returns:\n        A Plotly bar chart figure showing tomogram ID frequencies\n        \n    Example:\n        >>> fig = visualize_tomo_distribution(labels_df)\n        >>> fig.show(renderer='iframe')\n    \"\"\"\n    # (1) Count the frequency of each tomogram ID\n    tomo_counts = df['tomo_id'].value_counts().reset_index()\n    tomo_counts.columns = ['tomo_id', 'count']\n    \n    # (2) Sort by count in descending order\n    tomo_counts = tomo_counts.sort_values('count', ascending=False)\n    \n    # (3) Subset to only include the top n_top tomograms\n    top_tomos = tomo_counts.head(n_top)\n    \n    # (4) Create a horizontal bar chart for better readability of tomogram IDs\n    fig = px.bar(\n        top_tomos, \n        y='tomo_id', \n        x='count',\n        orientation='h',\n        color_discrete_sequence=[color],\n        title=f'<b>Top {n_top} Tomograms by Number of Motors</b>',  # Bold title\n        labels={'count': 'Number of Motors', 'tomo_id': 'Tomogram ID'},\n        height=height,\n        width=width\n    )\n    \n    # (5) Improve layout for better readability\n    fig.update_layout(\n        yaxis={'categoryorder': 'total ascending'},\n        xaxis_title='<b>Number of Motors</b>',  # Bold axis title\n        yaxis_title='<b>Tomogram ID</b>',  # Bold axis title\n        margin=dict(l=40, r=40, t=80, b=40),  # Increased margins for better spacing\n        title=dict(\n            text=f'<b>Top {n_top} Tomograms by Number of Motors</b>',\n            font=dict(size=22)  # Larger title\n        ),\n        title_x=0.5,  # Center the title\n        title_y=0.95,  # Position title higher\n        font=dict(family=\"Arial, sans-serif\"),  # Consistent font family\n    )\n    \n    # (6) Add data labels on bars and customize hover information\n    fig.update_traces(\n        texttemplate='<b>%{x}</b>',  # Bold text showing count\n        textposition='outside',  # Position text outside bars\n        textfont=dict(size=12, color=\"black\"),  # Text formatting\n        hovertemplate='<b>Tomogram ID:</b> %{y}<br><b>Number of Motors:</b> %{x}<extra></extra>'\n    )\n    \n    # (7) Apply consistent styling to axes\n    fig.update_xaxes(\n        showgrid=True,\n        gridwidth=1,\n        gridcolor='lightgray',\n        zeroline=True,\n        zerolinewidth=1,\n        zerolinecolor='black',\n        tickfont=dict(size=12),\n    )\n    \n    fig.update_yaxes(\n        tickfont=dict(size=12),\n        tickmode='linear'  # Ensure all ticks are shown\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T20:28:20.846719Z","iopub.execute_input":"2025-03-13T20:28:20.847067Z","iopub.status.idle":"2025-03-13T20:28:20.927448Z","shell.execute_reply.started":"2025-03-13T20:28:20.847040Z","shell.execute_reply":"2025-03-13T20:28:20.926479Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tomo_distribution_fig = visualize_tomo_distribution(labels_df, 25)\ntomo_distribution_fig.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">5.3 <b><code>Motor Axis [0,1,2]</code></b> EXPLORATION</h3>\n<hr><br>\n\n* `Motor axis 0` - the z-coordinate of the motor, i.e., which slice it is located on\n* `Motor axis 1` - the y-coordinate of the motor\n* `Motor axis 2` - the x-coordinate of the motor\n\n\n<br>","metadata":{}},{"cell_type":"code","source":"def visualize_motor_axis_distribution(\n    df: pd.DataFrame,\n    axis: int,\n    color: str = '#d27582',\n    bin_width: int | None = None,\n    height: int = 500,\n    width: int = 800\n) -> go.Figure:\n    \"\"\"Creates a histogram visualization showing the distribution of motor positions.\n    \n    Must be done along a specified axis (0: z, 1: y, 2: x).\n    \n    Args:\n        df (pd.DataFrame): \n            The tomogram dataset with at least 'Motor axis 0', 'Motor axis 1', 'Motor axis 2' columns\n        axis (int): \n            The axis to visualize (0 for z, 1 for y, 2 for x)\n        color (str, optional):\n            Color for the histogram bars (default: '#d27582')\n        bin_width (int, optional):\n            Width of histogram bins. If None, automatically determined (default: None)\n        height (int, optional): \n            Height of the figure in pixels (default: 500)\n        width (int, optional): \n            Width of the figure in pixels (default: 800)\n        \n    Returns:\n        A Plotly histogram figure showing distribution of motor positions along the specified axis\n        \n    Example:\n        >>> fig = visualize_motor_axis_distribution(labels_df, axis=0)\n        >>> fig.show(renderer='iframe')\n    \"\"\"\n    # Filter out rows with -1 values (no motor present)\n    filtered_df = df[df[f'Motor axis {axis}'] >= 0]\n    \n    # Define axis labels\n    axis_names = {0: 'Z (Slice)', 1: 'Y', 2: 'X'}\n    axis_label = axis_names[axis]\n    \n    # Determine bin width if not specified\n    if bin_width is None:\n        range_of_values = filtered_df[f'Motor axis {axis}'].max() - filtered_df[f'Motor axis {axis}'].min()\n        bin_width = max(1, round(range_of_values / 30))  # Default to 30 bins, minimum width of 1\n    \n    # Create histogram\n    fig = px.histogram(\n        filtered_df, \n        x=f'Motor axis {axis}',\n        nbins=None,  # Let Plotly determine bins based on bin_width\n        histnorm=None,  # Count values directly\n        color_discrete_sequence=[color],\n        title=f'<b>Distribution of Motor Positions - {axis_label} Axis</b>',\n        labels={f'Motor axis {axis}': f'{axis_label} Position'},\n        height=height,\n        width=width\n    )\n    \n    # Update layout for better readability\n    fig.update_layout(\n        xaxis_title=f'<b>{axis_label} Position</b>',\n        yaxis_title='<b>Count</b>',\n        margin=dict(l=40, r=40, t=80, b=40),\n        title=dict(\n            font=dict(size=22),  # Larger title\n        ),\n        title_x=0.5,  # Center the title\n        title_y=0.95,  # Position title higher\n        font=dict(family=\"Arial, sans-serif\"),\n        bargap=0.1  # Gap between bars\n    )\n    \n    # Customize axis appearance\n    fig.update_xaxes(\n        showgrid=True,\n        gridwidth=1,\n        gridcolor='lightgray',\n        zeroline=True,\n        zerolinewidth=1,\n        zerolinecolor='black',\n        tickfont=dict(size=12)\n    )\n    \n    fig.update_yaxes(\n        showgrid=True,\n        gridwidth=1,\n        gridcolor='lightgray',\n        zeroline=True,\n        zerolinewidth=1,\n        zerolinecolor='black',\n        tickfont=dict(size=12)\n    )\n    \n    # Enhance hover information\n    fig.update_traces(\n        hovertemplate=f'<b>{axis_label} Position</b>: %{{x}}<br><b>Count</b>: %{{y}}<extra></extra>'\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T20:32:25.617630Z","iopub.execute_input":"2025-03-13T20:32:25.618037Z","iopub.status.idle":"2025-03-13T20:32:25.627800Z","shell.execute_reply.started":"2025-03-13T20:32:25.618009Z","shell.execute_reply":"2025-03-13T20:32:25.626309Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize the Z-axis (slice) distribution\nz_dist = visualize_motor_axis_distribution(labels_df, axis=0)\nz_dist.show(renderer='iframe')\n\n# Visualize the Y-axis distribution\ny_dist = visualize_motor_axis_distribution(labels_df, axis=1)\ny_dist.show(renderer='iframe')\n\n# Visualize the X-axis distribution\nx_dist = visualize_motor_axis_distribution(labels_df, axis=2)\nx_dist.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_motor_3d_distribution(\n    df: pd.DataFrame,\n    color_by: str = 'tomogram_group',\n    opacity: float = 0.7,\n    marker_size: int = 5,\n    height: int = 700,\n    width: int = 900\n) -> go.Figure:\n    \"\"\"Creates a 3D scatter plot visualization.\n    \n    This will show the distribution of motor positions in 3D space (x, y, z coordinates).\n    \n    Args:\n        df (pd.DataFrame): \n            The tomogram dataset with at least 'Motor axis 0', 'Motor axis 1', 'Motor axis 2' columns\n        color_by (str, optional):\n            Attribute to use for coloring points. Options:\n            - 'tomogram_group': Group tomograms by motor count (1, 2-3, 4+)\n            - 'z_position': Color based on Z-axis position (slice number)\n            - 'voxel_spacing': Color based on tomogram resolution\n            - 'tomo_id': Color based on tomogram ID (not recommended for many tomograms)\n            - 'motor_count': Original coloring by number of motors in tomogram\n            (default: 'tomogram_group')\n        opacity (float, optional):\n            Opacity of the markers (default: 0.7)\n        marker_size (int, optional):\n            Size of the markers (default: 5)\n        height (int, optional): \n            Height of the figure in pixels (default: 700)\n        width (int, optional): \n            Width of the figure in pixels (default: 900)\n        \n    Returns:\n        A Plotly 3D scatter plot figure showing distribution of motor positions in 3D space\n        \n    Example:\n        >>> fig = visualize_motor_3d_distribution(labels_df, color_by='z_position')\n        >>> fig.show(renderer='iframe')\n    \"\"\"\n    # Filter out rows with -1 values (no motor present)\n    filtered_df = df[(df['Motor axis 0'] >= 0) & \n                     (df['Motor axis 1'] >= 0) & \n                     (df['Motor axis 2'] >= 0)].copy()\n    \n    # Create figure\n    fig = go.Figure()\n    \n    if color_by == 'tomogram_group':\n        # Create groups based on number of motors (1, 2-3, 4+)\n        filtered_df['group'] = pd.cut(\n            filtered_df['Number of motors'], \n            bins=[0, 1, 3, float('inf')],\n            labels=['Single Motor', '2-3 Motors', '4+ Motors']\n        )\n        \n        # Define colors for each group\n        colors = {\n            'Single Motor': 'rgb(31,119,180)',  # Blue\n            '2-3 Motors': 'rgb(255,127,14)',    # Orange\n            '4+ Motors': 'rgb(214,39,40)'       # Red\n        }\n        \n        # Plot each group separately\n        for group, color in colors.items():\n            group_df = filtered_df[filtered_df['group'] == group]\n            \n            if len(group_df) > 0:\n                hover_text = []\n                for idx, row in group_df.iterrows():\n                    hover_text.append(\n                        f\"<b>Tomogram ID</b>: {row['tomo_id']}<br>\"\n                        f\"<b>X Position</b>: {row['Motor axis 2']}<br>\"\n                        f\"<b>Y Position</b>: {row['Motor axis 1']}<br>\"\n                        f\"<b>Z Position</b>: {row['Motor axis 0']}<br>\"\n                        f\"<b>Group</b>: {group}<br>\"\n                        f\"<b>Motors in Tomogram</b>: {row['Number of motors']}\"\n                    )\n                \n                fig.add_trace(go.Scatter3d(\n                    x=group_df['Motor axis 2'],\n                    y=group_df['Motor axis 1'],\n                    z=group_df['Motor axis 0'],\n                    mode='markers',\n                    marker=dict(\n                        size=marker_size,\n                        color=color,\n                        opacity=opacity\n                    ),\n                    text=hover_text,\n                    hovertemplate=\"%{text}<extra></extra>\",\n                    name=group,\n                    showlegend=True\n                ))\n                \n    else:\n        # Prepare coloring based on selected attribute\n        color_data = None\n        colorscale = 'agsunset_r'\n        colorbar_title = \"\"\n        \n        if color_by == 'z_position':\n            color_data = filtered_df['Motor axis 0']\n            colorbar_title = \"<b>Z Position<br>(Slice Number)</b>\"\n            \n        elif color_by == 'voxel_spacing':\n            color_data = filtered_df['Voxel spacing']\n            colorbar_title = \"<b>Voxel Spacing<br>(Angstroms per Voxel)</b>\"\n            \n        elif color_by == 'tomo_id':\n            # Not recommended for many tomograms\n            filtered_df['tomo_id_code'] = pd.Categorical(filtered_df['tomo_id']).codes\n            color_data = filtered_df['tomo_id_code']\n            colorbar_title = \"<b>Tomogram ID</b>\"\n            \n        elif color_by == 'motor_count':\n            # Original coloring method\n            color_data = filtered_df['Number of motors']\n            colorbar_title = \"<b>Number of Motors<br>in Tomogram</b>\"\n        \n        else:\n            # Default to z-position if invalid option\n            color_data = filtered_df['Motor axis 0']\n            colorbar_title = \"<b>Z Position<br>(Slice Number)</b>\"\n        \n        # Add specific hover text\n        hover_text = []\n        for idx, row in filtered_df.iterrows():\n            hover_text.append(\n                f\"<b>Tomogram ID</b>: {row['tomo_id']}<br>\"\n                f\"<b>X Position</b>: {row['Motor axis 2']}<br>\"\n                f\"<b>Y Position</b>: {row['Motor axis 1']}<br>\"\n                f\"<b>Z Position</b>: {row['Motor axis 0']}<br>\"\n                f\"<b>Motors in Tomogram</b>: {row['Number of motors']}<br>\"\n                f\"<b>Voxel Spacing</b>: {row['Voxel spacing']}\"\n            )\n        \n        # Create 3D scatter plot with continuous color scale\n        fig.add_trace(go.Scatter3d(\n            x=filtered_df['Motor axis 2'],\n            y=filtered_df['Motor axis 1'],\n            z=filtered_df['Motor axis 0'],\n            mode='markers',\n            marker=dict(\n                size=marker_size,\n                color=color_data,\n                colorscale=colorscale,\n                opacity=opacity,\n                colorbar=dict(\n                    title=colorbar_title,\n                    thickness=20,\n                    x=0.9\n                )\n            ),\n            text=hover_text,\n            hovertemplate=\"%{text}<extra></extra>\",\n            showlegend=False\n        ))\n    \n    # Determine title based on coloring method\n    title_text = '<b>3D Distribution of Motor Positions</b>'\n    if color_by == 'tomogram_group':\n        title_text += '<br><sup>Colored by groups: Single Motor, 2-3 Motors, 4+ Motors</sup>'\n    elif color_by == 'z_position':\n        title_text += '<br><sup>Colored by Z Position (Slice Number)</sup>'\n    elif color_by == 'voxel_spacing':\n        title_text += '<br><sup>Colored by Voxel Spacing (Resolution)</sup>'\n    elif color_by == 'tomo_id':\n        title_text += '<br><sup>Colored by Tomogram ID</sup>'\n    elif color_by == 'motor_count':\n        title_text += '<br><sup>Colored by Number of Motors in Tomogram</sup>'\n    \n    # Update layout for better readability\n    fig.update_layout(\n        title=dict(\n            text=title_text,\n            font=dict(size=22)\n        ),\n        scene=dict(\n            xaxis_title='<b>X Position</b>',\n            yaxis_title='<b>Y Position</b>',\n            zaxis_title='<b>Z Position</b>',\n            xaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n            yaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n            zaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n        ),\n        margin=dict(l=0, r=0, t=100, b=0),  # Increased top margin for subtitle\n        title_x=0.5,\n        title_y=0.97,\n        height=height,\n        width=width,\n        font=dict(family=\"Arial, sans-serif\"),\n        legend=dict(\n            title=\"<b>Tomogram Group</b>\",\n            yanchor=\"top\",\n            y=0.99,\n            xanchor=\"left\",\n            x=0.01,\n            bgcolor=\"rgba(255, 255, 255, 0.6)\",\n            bordercolor=\"gray\",\n            borderwidth=1,\n            itemsizing='constant'\n        )\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T20:44:19.041980Z","iopub.execute_input":"2025-03-13T20:44:19.042327Z","iopub.status.idle":"2025-03-13T20:44:19.137410Z","shell.execute_reply.started":"2025-03-13T20:44:19.042301Z","shell.execute_reply":"2025-03-13T20:44:19.136264Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Default - Group by motor count (1, 2-3, 4+)\nfig1 = visualize_motor_3d_distribution(labels_df, color_by=\"tomogram_group\")\nfig1.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Color by Z-position (depth/slice)\nfig2 = visualize_motor_3d_distribution(labels_df, color_by='z_position')\nfig2.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T20:44:20.606814Z","iopub.execute_input":"2025-03-13T20:44:20.607170Z","iopub.status.idle":"2025-03-13T20:44:20.679611Z","shell.execute_reply.started":"2025-03-13T20:44:20.607143Z","shell.execute_reply":"2025-03-13T20:44:20.678535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Color by tomogram resolution\nfig3 = visualize_motor_3d_distribution(labels_df, color_by='voxel_spacing')\nfig3.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T20:44:20.785546Z","iopub.execute_input":"2025-03-13T20:44:20.785944Z","iopub.status.idle":"2025-03-13T20:44:20.860838Z","shell.execute_reply.started":"2025-03-13T20:44:20.785915Z","shell.execute_reply":"2025-03-13T20:44:20.859618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_motor_positions_by_tomogram(\n    df: pd.DataFrame,\n    tomogram_id_list: list[str] | None = None,\n    n_top: int = 5,\n    height: int = 500,\n    width: int = 800\n) -> dict[str, go.Figure]:\n    \"\"\"Creates a set of visualizations showing motor positions within the top N tomograms.\n\n    We default to using the tomograms with the most motors to make things easier if no list is provided.\n    \n    Args:\n        df (pd.DataFrame): \n            The tomogram dataset with at least 'tomo_id', 'Motor axis 0/1/2' columns\n        tomogram_id_list (list[str], optional):\n            Optional list of tomograms to return as plottable figures in dictionary\n        n_top (int, optional):\n            Number of top tomograms to visualize (default: 10)\n        height (int, optional): \n            Height of each figure in pixels (default: 500)\n        width (int, optional): \n            Width of each figure in pixels (default: 800)\n        \n    Returns:\n        A dictionary of Plotly figures showing motor positions within each tomogram\n        \n    Example:\n        >>> figs = visualize_motor_positions_by_tomogram(labels_df, n_top=5)\n        >>> for tomo_id, fig in figs.items():\n        >>>     fig.show(renderer='iframe')\n    \"\"\"\n    if not tomogram_id_list:\n        # Filter out rows with -1 values (no motor present)\n        filtered_df = df[(df['Motor axis 0'] >= 0) & \n                         (df['Motor axis 1'] >= 0) & \n                         (df['Motor axis 2'] >= 0)]\n    else:\n        filtered_df = df[df['tomo_id'].isin(tomogram_id_list)]\n    \n    # Get top N tomograms with most motors\n    top_tomos = filtered_df['tomo_id'].value_counts().head(n_top).index.tolist()\n    \n    # Create a figure for each top tomogram\n    figures = {}\n    \n    for tomo_id in top_tomos:\n        tomo_df = filtered_df[filtered_df['tomo_id'] == tomo_id]\n        \n        # Get the dimensions of this tomogram\n        tomo_shape = (\n            tomo_df['Array shape (axis 0)'].iloc[0],\n            tomo_df['Array shape (axis 1)'].iloc[0],\n            tomo_df['Array shape (axis 2)'].iloc[0]\n        )\n        \n        motor_count = tomo_df['Number of motors'].iloc[0]\n        if motor_count==0:\n            print(f\"\\n... [SKIPPING] No Motors Found For tomo_id={tomo_id} [SKIPPING] ...\\n\")\n            continue\n            \n        # Create 3D scatter plot for this tomogram\n        fig = go.Figure(data=[go.Scatter3d(\n            x=tomo_df['Motor axis 2'],\n            y=tomo_df['Motor axis 1'],\n            z=tomo_df['Motor axis 0'],\n            mode='markers',\n            marker=dict(\n                size=10,\n                color='red',\n                symbol='circle',  # Valid symbol for Scatter3d\n                opacity=0.8\n            ),\n            hovertemplate=\"<b>Motor Position</b><br>\" +\n                          \"X: %{x}<br>\" +\n                          \"Y: %{y}<br>\" +\n                          \"Z: %{z}<extra></extra>\",\n            name=\"Motors\"\n        )])\n        \n        # Create wireframe box to represent tomogram boundaries\n        fig = add_wireframe_box(\n            fig, \n            x0=0, y0=0, z0=0, \n            x1=tomo_shape[2], y1=tomo_shape[1], z1=tomo_shape[0]\n        )\n        \n        # Update layout for better readability\n        fig.update_layout(\n            title=dict(\n                text=f'<b>Motor Positions in {tomo_id}</b><br><sup>Total Motors: {motor_count}</sup>',\n                font=dict(size=18)\n            ),\n            scene=dict(\n                xaxis_title='<b>X Position</b>',\n                yaxis_title='<b>Y Position</b>',\n                zaxis_title='<b>Z Position</b>',\n                aspectmode='data',  # Preserve the shape proportions\n                camera=dict(\n                    eye=dict(x=1.5, y=1.5, z=1.5)  # Adjust camera position for better view\n                )\n            ),\n            margin=dict(l=0, r=0, t=80, b=0),\n            title_x=0.5,\n            title_y=0.97,\n            height=height,\n            width=width,\n            font=dict(family=\"Arial, sans-serif\"),\n            showlegend=True,\n            legend=dict(\n                title=\"<b>Components</b>\",\n                yanchor=\"top\",\n                y=0.99,\n                xanchor=\"left\",\n                x=0.01,\n                bgcolor=\"rgba(255, 255, 255, 0.6)\",\n                bordercolor=\"gray\",\n                borderwidth=1\n            )\n        )\n        \n        figures[tomo_id] = fig\n    \n    return figures\n\n\ndef add_wireframe_box(fig: go.Figure, x0: int, y0: int, z0: int, x1: int, y1: int, z1: int) -> go.Figure:\n    \"\"\"Adds a wireframe box to a 3D figure to represent tomogram boundaries.\n    \n    Args:\n        fig (go.Figure): Plotly figure object\n        x0 (int): Minimum coordinates (x position)\n        y0 (int): Minimum coordinates (y position)\n        z0 (int): Minimum coordinates (z position)\n        x1 (int): Maximum coordinates (x position)\n        y1 (int): Maximum coordinates (y position)\n        z1 (int): Maximum coordinates (z position)\n        \n    Returns:\n        go.Figure:\n            Updated Plotly figure with wireframe box\n    \"\"\"\n    # Create the 8 corners of the box\n    x = [x0, x1, x1, x0, x0, x1, x1, x0]\n    y = [y0, y0, y1, y1, y0, y0, y1, y1]\n    z = [z0, z0, z0, z0, z1, z1, z1, z1]\n    \n    # Define the 12 lines connecting the corners\n    lines = [\n        # Bottom face\n        [0, 1], [1, 2], [2, 3], [3, 0],\n        # Top face\n        [4, 5], [5, 6], [6, 7], [7, 4],\n        # Connecting edges\n        [0, 4], [1, 5], [2, 6], [3, 7]\n    ]\n    \n    # Add each line as a separate trace\n    for line in lines:\n        fig.add_trace(go.Scatter3d(\n            x=[x[line[0]], x[line[1]]],\n            y=[y[line[0]], y[line[1]]],\n            z=[z[line[0]], z[line[1]]],\n            mode='lines',\n            line=dict(color='blue', width=2),\n            hoverinfo='none',\n            showlegend=False\n        ))\n    \n    # Add a helper trace for the legend\n    fig.add_trace(go.Scatter3d(\n        x=[None], y=[None], z=[None],\n        mode='lines',\n        line=dict(color='blue', width=2),\n        name='Tomogram Boundary',\n        showlegend=True\n    ))\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:02:31.832767Z","iopub.execute_input":"2025-03-13T21:02:31.833117Z","iopub.status.idle":"2025-03-13T21:02:31.888079Z","shell.execute_reply.started":"2025-03-13T21:02:31.833091Z","shell.execute_reply":"2025-03-13T21:02:31.887091Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate visualizations for top 5 tomograms with most motors\n# tomo_figs = visualize_motor_positions_by_tomogram(labels_df, n_top=5)\ntomo_figs = visualize_motor_positions_by_tomogram(labels_df, tomogram_id_list=['tomo_226cd8', 'tomo_003acc'])\n\n# Display each figure individually\nfor tomo_id, fig in tomo_figs.items():\n    rich.print(f\"\\n\\n[bold red]Displaying visualization for tomogram: {tomo_id}[/bold red]\\n\")\n    fig.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">5.4 <b><code>Array shape Axis [0,1,2]</code></b> EXPLORATION</h3>\n<hr><br>\n\n* `Array shape axis 0` - z-axis length, i.e., number of slices in the tomogram\n* `Array shape axis 1` - y-axis length, or width of each slice\n* `Array shape axis 2` - x-axis length, or height of each slice\n\n<br>","metadata":{}},{"cell_type":"code","source":"# def visualize_tomogram_dimensions(\n#     df: pd.DataFrame,\n#     height: int = 500,\n#     width: int = 800\n# ) -> go.Figure:\n#     \"\"\"\n#     Creates a 3D scatter plot visualizing tomogram dimensions with each point \n#     representing a unique tomogram.\n    \n#     Args:\n#         df (pd.DataFrame): \n#             The tomogram dataset with 'Array shape' columns\n#         height (int, optional): \n#             Height of the figure in pixels (default: 500)\n#         width (int, optional): \n#             Width of the figure in pixels (default: 800)\n        \n#     Returns:\n#         A Plotly 3D scatter plot showing tomogram dimensions\n        \n#     Example:\n#         >>> fig = visualize_tomogram_dimensions(labels_df)\n#         >>> fig.show(renderer='iframe')\n#     \"\"\"\n#     # Get unique tomograms\n#     unique_tomos = df.drop_duplicates(subset=['tomo_id'])\n    \n#     # Create hover text\n#     hover_text = []\n#     for _, row in unique_tomos.iterrows():\n#         hover_text.append(\n#             f\"<b>Tomogram ID</b>: {row['tomo_id']}<br>\" +\n#             f\"<b>Z Dimension</b>: {row['Array shape (axis 0)']}<br>\" +\n#             f\"<b>Y Dimension</b>: {row['Array shape (axis 1)']}<br>\" +\n#             f\"<b>X Dimension</b>: {row['Array shape (axis 2)']}<br>\" +\n#             f\"<b>Voxel Spacing</b>: {row['Voxel spacing']}<br>\" +\n#             f\"<b>Number of Motors</b>: {row['Number of motors']}\"\n#         )\n    \n#     # Create 3D scatter plot\n#     fig = go.Figure(data=[go.Scatter3d(\n#         x=unique_tomos['Array shape (axis 2)'],  # X dimension\n#         y=unique_tomos['Array shape (axis 1)'],  # Y dimension\n#         z=unique_tomos['Array shape (axis 0)'],  # Z dimension\n#         mode='markers',\n#         marker=dict(\n#             size=8,\n#             color=unique_tomos['Number of motors'],\n#             colorscale='Viridis',\n#             opacity=0.8,\n#             colorbar=dict(\n#                 title=\"<b>Number of Motors</b>\",\n#                 thickness=20,\n#             )\n#         ),\n#         text=hover_text,\n#         hovertemplate=\"%{text}<extra></extra>\"\n#     )])\n    \n#     # Update layout\n#     fig.update_layout(\n#         title=dict(\n#             text='<b>Tomogram Dimensions</b><br><sup>Each point represents one tomogram</sup>',\n#             font=dict(size=22)\n#         ),\n#         scene=dict(\n#             xaxis_title='<b>X Dimension (pixels)</b>',\n#             yaxis_title='<b>Y Dimension (pixels)</b>',\n#             zaxis_title='<b>Z Dimension (slices)</b>',\n#             xaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n#             yaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n#             zaxis=dict(showgrid=True, gridwidth=1, gridcolor='lightgray'),\n#         ),\n#         margin=dict(l=0, r=0, t=100, b=0),\n#         title_x=0.5,\n#         title_y=0.97,\n#         height=height,\n#         width=width,\n#         font=dict(family=\"Arial, sans-serif\")\n#     )\n    \n#     return fig\n\n# # Visualize tomogram dimensions in 3D space\n# dim_fig = visualize_tomogram_dimensions(labels_df)\n# dim_fig.show(renderer='iframe')\n\n\ndef visualize_tomogram_dimension_distribution(\n    df: pd.DataFrame,\n    axis: int,\n    color: str = '#d27582',\n    bin_width: int | None = None,\n    height: int = 500,\n    width: int = 800\n) -> go.Figure:\n    \"\"\"Creates a histogram visualization showing the distribution of tomogram dimensions.\n    \n    Visualizes distribution along a specified axis (0: z, 1: y, 2: x).\n    \n    Args:\n        df (pd.DataFrame): \n            The tomogram dataset with 'Array shape' columns\n        axis (int): \n            The axis to visualize (0 for z, 1 for y, 2 for x)\n        color (str, optional):\n            Color for the histogram bars (default: '#2c7fb8')\n        bin_width (int, optional):\n            Width of histogram bins. If None, automatically determined (default: None)\n        height (int, optional): \n            Height of the figure in pixels (default: 500)\n        width (int, optional): \n            Width of the figure in pixels (default: 800)\n        \n    Returns:\n        A Plotly histogram figure showing distribution of tomogram dimensions along specified axis\n        \n    Example:\n        >>> fig = visualize_tomogram_dimension_distribution(labels_df, axis=0)\n        >>> fig.show(renderer='iframe')\n    \"\"\"\n    # Get unique tomograms\n    unique_tomos = df.drop_duplicates(subset=['tomo_id'])\n    \n    # Define axis labels and column names\n    axis_names = {0: 'Z (Slices)', 1: 'Y (Height)', 2: 'X (Width)'}\n    axis_label = axis_names[axis]\n    column_name = f'Array shape (axis {axis})'\n    \n    # Determine bin width if not specified\n    if bin_width is None:\n        range_of_values = unique_tomos[column_name].max() - unique_tomos[column_name].min()\n        bin_width = max(1, round(range_of_values / 25))  # Default to 25 bins, minimum width of 1\n    \n    # Create histogram\n    fig = px.histogram(\n        unique_tomos, \n        x=column_name,\n        nbins=None,  # Let Plotly determine bins based on bin_width\n        histnorm=None,  # Count values directly\n        color_discrete_sequence=[color],\n        title=f'<b>Distribution of Tomogram Dimensions - {axis_label}</b>',\n        labels={column_name: f'{axis_label} Dimension (pixels)'},\n        height=height,\n        width=width\n    )\n    \n    # Update layout for better readability\n    fig.update_layout(\n        xaxis_title=f'<b>{axis_label} Dimension (pixels)</b>',\n        yaxis_title='<b>Count</b>',\n        margin=dict(l=40, r=40, t=80, b=40),\n        title=dict(\n            font=dict(size=22),  # Larger title\n        ),\n        title_x=0.5,  # Center the title\n        title_y=0.95,  # Position title higher\n        font=dict(family=\"Arial, sans-serif\"),\n        bargap=0.1  # Gap between bars\n    )\n    \n    # Customize axis appearance\n    fig.update_xaxes(\n        showgrid=True,\n        gridwidth=1,\n        gridcolor='lightgray',\n        zeroline=True,\n        zerolinewidth=1,\n        zerolinecolor='black',\n        tickfont=dict(size=12)\n    )\n    \n    fig.update_yaxes(\n        showgrid=True,\n        gridwidth=1,\n        gridcolor='lightgray',\n        zeroline=True,\n        zerolinewidth=1,\n        zerolinecolor='black',\n        tickfont=dict(size=12)\n    )\n    \n    # Add mean line\n    mean_dimension = unique_tomos[column_name].mean()\n    fig.add_vline(\n        x=mean_dimension, \n        line_dash=\"dash\", \n        line_color=\"yellow\",\n        annotation_text=f\"Mean: {mean_dimension:.1f}\",\n        annotation_position=\"top right\"\n    )\n    \n    # Add median line\n    median_dimension = unique_tomos[column_name].median()\n    fig.add_vline(\n        x=median_dimension, \n        line_dash=\"dot\", \n        line_color=\"blue\",\n        annotation_text=f\"Median: {median_dimension:.1f}\",\n        annotation_position=\"top left\"\n    )\n    \n    # Enhance hover information\n    fig.update_traces(\n        hovertemplate=f'<b>{axis_label} Dimension</b>: %{{x}}<br><b>Count</b>: %{{y}}<extra></extra>'\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:11:42.195862Z","iopub.execute_input":"2025-03-13T21:11:42.196195Z","iopub.status.idle":"2025-03-13T21:11:42.435341Z","shell.execute_reply.started":"2025-03-13T21:11:42.196170Z","shell.execute_reply":"2025-03-13T21:11:42.434327Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize the Z-axis (slice) distribution\nz_dist_fig = visualize_tomogram_dimension_distribution(labels_df, axis=0)\nz_dist_fig.show(renderer='iframe')\n\n# Visualize the Y-axis distribution\ny_dist_fig = visualize_tomogram_dimension_distribution(labels_df, axis=1)\ny_dist_fig.show(renderer='iframe')\n\n# Visualize the X-axis distribution\nx_dist_fig = visualize_tomogram_dimension_distribution(labels_df, axis=2)\nx_dist_fig.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_tomogram_shape_profiles(\n    df: pd.DataFrame,\n    n_profiles: int = 10,\n    height: int = 600,\n    width: int = 900\n) -> go.Figure:\n    \"\"\"Creates a visualization showing dimension profiles of tomograms.\n    \n    This displays the relative dimensions across all three axes for the most common\n    tomogram shapes in the dataset.\n    \n    Args:\n        df (pd.DataFrame): \n            The tomogram dataset with 'Array shape' columns\n        n_profiles (int, optional):\n            Number of most common tomogram profiles to show (default: 10)\n        height (int, optional): \n            Height of the figure in pixels (default: 600)\n        width (int, optional): \n            Width of the figure in pixels (default: 900)\n        \n    Returns:\n        A Plotly figure showing the most common tomogram shape profiles\n        \n    Example:\n        >>> fig = visualize_tomogram_shape_profiles(labels_df)\n        >>> fig.show(renderer='iframe'\n    \"\"\"\n    # Get unique tomograms\n    unique_tomos = df.drop_duplicates(subset=['tomo_id']).copy()\n    \n    # Create a shape profile string for each tomogram\n    unique_tomos['shape_profile'] = unique_tomos.apply(\n        lambda row: f\"{int(row['Array shape (axis 0)'])}×{int(row['Array shape (axis 1)'])}×{int(row['Array shape (axis 2)'])}\",\n        axis=1\n    )\n    \n    # Get the most common profiles\n    top_profiles = unique_tomos['shape_profile'].value_counts().head(n_profiles)\n    profile_counts = top_profiles.reset_index()\n    profile_counts.columns = ['shape_profile', 'count']\n    \n    # Prepare data for radar/polar chart\n    shapes = []\n    for profile in top_profiles.index:\n        # Extract dimensions from profile string\n        dims = [int(d) for d in profile.split('×')]\n        \n        # Add to shapes list\n        shapes.append({\n            'shape_profile': profile,\n            'count': top_profiles[profile],\n            'Z': dims[0],\n            'Y': dims[1],\n            'X': dims[2]\n        })\n    \n    # Create DataFrame for plotting\n    shapes_df = pd.DataFrame(shapes)\n    \n    # Normalize dimensions for radar chart\n    for axis in ['Z', 'Y', 'X']:\n        max_val = shapes_df[axis].max()\n        shapes_df[f'{axis}_norm'] = shapes_df[axis] / max_val\n    \n    # Create figure\n    fig = go.Figure()\n    \n    # Add a trace for each profile\n    for i, row in shapes_df.iterrows():\n        fig.add_trace(go.Scatterpolar(\n            r=[row['Z_norm'], row['Y_norm'], row['X_norm'], row['Z_norm']],  # Close the loop\n            theta=['Z', 'Y', 'X', 'Z'],  # Close the loop\n            fill='toself',\n            name=f\"{row['shape_profile']} (n={row['count']})\",\n            hoverinfo='text',\n            hovertext=(\n                f\"<b>Profile</b>: {row['shape_profile']}<br>\"\n                f\"<b>Count</b>: {row['count']}<br>\"\n                f\"<b>Z</b>: {row['Z']}<br>\"\n                f\"<b>Y</b>: {row['Y']}<br>\"\n                f\"<b>X</b>: {row['X']}\"\n            )\n        ))\n    \n    # Update layout\n    fig.update_layout(\n        title=dict(\n            text=f'<b>Top {n_profiles} Tomogram Shape Profiles</b><br><sup>Normalized dimensions</sup>',\n            font=dict(size=22)\n        ),\n        polar=dict(\n            radialaxis=dict(\n                visible=True,\n                range=[0, 1]\n            )\n        ),\n        showlegend=True,\n        legend=dict(\n            title=\"<b>Shape Profiles</b>\",\n            yanchor=\"top\",\n            y=0.99,\n            xanchor=\"right\",\n            x=0.99,\n            bgcolor=\"rgba(255, 255, 255, 0.6)\",\n            bordercolor=\"gray\",\n            borderwidth=1\n        ),\n        margin=dict(l=80, r=80, t=100, b=80),\n        title_x=0.5,\n        title_y=0.95,\n        height=height,\n        width=width,\n        font=dict(family=\"Arial, sans-serif\")\n    )\n    \n    return fig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:14:11.456267Z","iopub.execute_input":"2025-03-13T21:14:11.456608Z","iopub.status.idle":"2025-03-13T21:14:11.506860Z","shell.execute_reply.started":"2025-03-13T21:14:11.456582Z","shell.execute_reply":"2025-03-13T21:14:11.505882Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"shape_profiles_fig = visualize_tomogram_shape_profiles(labels_df)\nshape_profiles_fig.show(renderer='iframe')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<h3 style=\"font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: #6c0f1c; background-color: #ffffff;\">5.5 <b>VISUAL</b> EXPLORATION</h3>\n<hr><br>\n\nLet's see 'em","metadata":{}},{"cell_type":"code","source":"def visualize_slice_with_motors(\n    tomo_id: str,\n    box_size: int = 40,\n    figsize: tuple = (20, 20),\n    slice_idx: int | None = None,\n    df: pd.DataFrame | None = None,\n    train_dir: str | None = None,\n    show_titlebox: bool = False\n) -> tuple:\n    \"\"\"Visualizes a tomogram slice with motors highlighted using colored transparent patches.\n    \n    This function creates a visualization of a single tomogram slice with annotated\n    motors using semi-transparent colored boxes. Each motor gets a distinct color\n    to improve visibility of small structures.\n    \n    Args:\n        tomo_id (str): String identifier for the tomogram.\n        box_size (int, optional): Integer size of bounding box to draw around motors (default: 24).\n        figsize (tuple[int], optional): Tuple of (width, height) for the output figure (default: (12, 10)).\n        slice_idx (int, optional): Integer Z-axis slice index to visualize. If unset, we will use the first slice with a motor.\n        df (pd.Dataframe, optional): Pandas DataFrame containing motor coordinates and metadata.\n        train_dir (str, optional): String path to directory containing tomogram data.\n        show_titlebox (bool, optional): Display optional bounding box label and instance count\n        \n    Returns:\n        tuple: (fig, ax) Matplotlib figure and axes objects for further customization.\n        \n    Raises:\n        FileNotFoundError: If the specified slice image cannot be found.\n    \"\"\"\n    # (0) Set defaults if not already set\n    df = df or labels_df\n    train_dir = train_dir or TRAIN_DIR\n    \n    # If no slice_idx is provided, use the first slice with a motor\n    if slice_idx is None:\n        motors_in_tomo = df[df['tomo_id'] == tomo_id]\n        if len(motors_in_tomo) == 0:\n            print(f\"Error: No motors found for tomogram {tomo_id}\")\n            return None, None\n        slice_idx = int(motors_in_tomo[\"Motor axis 0\"].values[0])\n    \n    # (1) Construct the path to the slice image file\n    slice_path = os.path.join(train_dir, tomo_id, f\"slice_{slice_idx:04d}.jpg\")\n    \n    # (2) Verify the slice exists and load it\n    if not os.path.exists(slice_path):\n        print(f\"Error: Slice {slice_path} does not exist\")\n        return None, None\n    \n    # (3) Load the image data as a numpy array\n    img = np.array(Image.open(slice_path))\n    \n    # (4) Filter the dataframe to get motors for this specific tomogram and slice\n    motors = df[(df['tomo_id'] == tomo_id) & (df['Motor axis 0'] == slice_idx)]\n    \n    # (5) Create a new figure and axis for visualization\n    fig, ax = plt.subplots(figsize=figsize)\n    \n    # (6) Display the grayscale tomogram slice\n    ax.imshow(img, cmap='gray')\n    \n    # (7) Define a carefully selected color palette for motor annotations\n    # These colors are chosen to be visually distinct but harmonious\n    motor_colors = [\n        '#1f77b4',  # Blue\n        '#ff7f0e',  # Orange\n        '#2ca02c',  # Green\n        '#d62728',  # Red\n        '#9467bd',  # Purple\n        '#8c564b',  # Brown\n        '#e377c2',  # Pink\n        '#7f7f7f',  # Gray\n        '#bcbd22',  # Olive\n        '#17becf'   # Cyan\n    ]\n    \n    # (8) Annotate each motor in the current slice\n    for i, (_, motor) in enumerate(motors.iterrows()):\n        # (8a) Select a color for this motor, cycling through the palette if needed\n        color_idx = i % len(motor_colors)\n        motor_color = motor_colors[color_idx]\n        \n        # (8b) Extract the motor coordinates (integer pixel positions)\n        y = int(motor['Motor axis 1'])\n        x = int(motor['Motor axis 2'])\n        half_box = box_size // 2\n        \n        # (8c) Create a semi-transparent rectangle to highlight the motor\n        # Thin border with matching fill color for optimal visibility\n        rect = Rectangle(\n            (x - half_box, y - half_box),  # Upper-left corner position\n            box_size, box_size,            # Width and height\n            linewidth=3.0,                 # Border\n            edgecolor=motor_color,         # Border color\n            facecolor=motor_color,         # Fill with same color as border\n            alpha=0.3                      # Semi-transparent fill for visibility\n        )\n        ax.add_patch(rect)\n        \n        # (8d) Add a label identifying the motor with good contrast\n        if show_titlebox:\n            ax.text(\n                x, y - half_box - 5,           # Position just above the box\n                f\"Motor {i+1}\",                # Label text with motor number\n                color='white',                 # White text for readability\n                fontsize=9,                    # Readable but not oversized font\n                fontweight='bold',             # Bold for visibility against background\n                bbox=dict(                     # Background box for contrast\n                    facecolor=motor_color,     # Same color as the motor annotation\n                    alpha=0.8,                 # Mostly opaque for readability\n                    pad=1,                     # Small padding around text\n                    boxstyle='round,pad=0.3'   # Slightly rounded corners\n                ),\n                ha='center',                   # Center-align text horizontally\n                zorder=10                      # Ensure text appears above other elements\n            )\n    \n    # (9) Add informative title and metadata\n    ax.set_title(f\"Tomogram: {tomo_id}, Slice: {slice_idx}\", fontsize=14, fontweight='bold')\n    \n    # (10) Add a count of motors in the current slice\n    ax.text(\n        10, 20,                            # Position in upper-left corner\n        f\"Total Motors: {len(motors)}\",    # Display count of motors\n        color='white',                     # White text for visibility\n        bbox=dict(                         # Background box\n            facecolor='black',             # Black background\n            alpha=0.5,                     # Partially transparent\n            boxstyle='round,pad=0.5'       # Rounded corners\n        ), \n        fontsize=10\n    )\n    \n    # (11) Add a legend for the motor identifiers if motors are present\n    if len(motors) > 0:\n        # (11a) Create legend elements for each motor\n        legend_elements = []\n        for i in range(min(len(motors), len(motor_colors))):\n            color = motor_colors[i % len(motor_colors)]\n            legend_elements.append(\n                plt.Line2D(\n                    [0], [0],                  # Dummy coordinates\n                    marker='s',                # Square marker matching annotations\n                    color='w',                 # White edge\n                    markerfacecolor=color,     # Fill with motor color\n                    markersize=8,              # Visible but not too large\n                    label=f'Motor {i+1}'       # Label with motor number\n                )\n            )\n        \n        # (11b) Place the legend in the upper right corner\n        ax.legend(\n            handles=legend_elements,\n            loc='upper right',\n            title='Motors',\n            framealpha=0.7,                # Semi-transparent background\n            fontsize='small',\n            title_fontsize='small'\n        )\n    \n    # (12) Hide axis for a cleaner visualization\n    ax.axis('off')\n    \n    # (13) Ensure layout is clean and tight\n    plt.tight_layout()\n    \n    # (14) Return the figure and axes for potential further customization\n    return fig, ax\n\n\ndef visualize_multiple_slices(\n    tomo_id: str,\n    start_slice: int | None = None,\n    end_slice: int | None = None, \n    step: int = 1,\n    box_size: int = 30,\n    figsize_x: int = 20,\n    df: pd.DataFrame | None = None,\n    train_dir: str | None = None,\n    show_titlebox: bool = False\n) -> tuple:\n    \"\"\"Visualizes multiple tomogram slices with motors highlighted.\n    \n    Creates a grid of visualizations showing multiple consecutive slices from a \n    tomogram to help understand the 3D distribution of motors. Consistent colors\n    are used across slices for better tracking of structures. Only shows slices\n    that contain motors.\n    \n    Args:\n        tomo_id (str): String identifier for the tomogram.\n        start_slice (int, optional): Integer starting Z-axis slice index. If None, uses first slice with a motor.\n        end_slice (int, optional): Integer ending Z-axis slice index. If None, uses last slice with a motor.\n        step (int, optional): Integer step size between visualized slices. Default is 1 (show all motor slices).\n        box_size (int, optional): Integer size of bounding box to draw around motors (default: 34).\n        figsize_x (int, optional): X dimension for output figure\n        df (pd.Dataframe, optional): Pandas DataFrame containing motor coordinates and metadata.\n        train_dir (str, optional): String path to directory containing tomogram data.\n        show_titlebox (bool, optional): Display optional bounding box label and instance count\n        \n    Returns:\n        tuple: (fig, axes) Matplotlib figure and axes objects for further customization.\n    \"\"\"\n    # (0) Set defaults if not already set\n    # Use the global labels_df if no dataframe is provided\n    df = df or labels_df\n    # Use the global TRAIN_DIR if no directory is provided\n    train_dir = train_dir or TRAIN_DIR\n    \n    # (1) Get all slice indices that contain motors for this tomogram\n    # First, filter the dataframe to only include motors in the requested tomogram\n    motors_in_tomo = df[df['tomo_id'] == tomo_id]\n    \n    # Check if any motors exist for this tomogram\n    if len(motors_in_tomo) == 0:\n        print(f\"Error: No motors found for tomogram {tomo_id}\")\n        return None, None\n    \n    # Extract unique slice indices that contain motors, convert to int, and sort\n    # This ensures we only show slices that actually have motors\n    all_motor_slices = sorted(motors_in_tomo[\"Motor axis 0\"].astype(int).unique().tolist())\n    \n    # (2) Apply start_slice and end_slice filters if provided\n    # If start_slice is specified, only include slices at or after that index\n    if start_slice is not None:\n        all_motor_slices = [s for s in all_motor_slices if s >= start_slice]\n    \n    # If end_slice is specified, only include slices at or before that index\n    if end_slice is not None:\n        all_motor_slices = [s for s in all_motor_slices if s <= end_slice]\n    \n    # (3) Apply step to select slices\n    # Using Python's slice notation to take every 'step' slice\n    # For example, step=2 will show every other slice\n    slices = all_motor_slices[::step]\n    \n    # After applying all filters, verify we still have slices to display\n    if len(slices) == 0:\n        print(f\"Error: No motor slices found for tomogram {tomo_id} with the given parameters\")\n        return None, None\n    \n    # Store the number of slices for later use in grid calculations\n    n_slices = len(slices)\n    \n    # (4) Calculate optimal grid dimensions for the subplots\n    # Limit to 3 columns maximum for readability\n    cols = min(3, n_slices)  \n    # Calculate how many rows are needed to fit all slices\n    # Using integer division with ceiling to ensure all slices fit\n    rows = (n_slices + cols - 1) // cols\n    \n    # (5) Create a figure with a grid of subplots\n    # This creates a single figure with an array of axes (subplots)\n    figsize_y = 8*rows\n    fig, axes = plt.subplots(rows, cols, figsize=(figsize_x, figsize_y))\n    \n    # (6) Handle different axes array shapes based on grid dimensions\n    # Matplotlib returns different structures depending on grid shape:\n    if rows == 1 and cols == 1:\n        # For a single subplot, convert to a 2D array for consistent indexing\n        axes = np.array([[axes]])\n    elif rows == 1:\n        # For a single row, reshape to 2D array with shape (1, cols)\n        axes = axes.reshape(1, -1)\n    elif cols == 1:\n        # For a single column, reshape to 2D array with shape (rows, 1)\n        axes = axes.reshape(-1, 1)\n    \n    # (7) Define a consistent color palette for motor annotations\n    # These colors are chosen to be distinct but visually harmonious\n    # Using the Matplotlib default color cycle for consistency\n    motor_colors = [\n        '#1f77b4',  # Blue\n        '#ff7f0e',  # Orange\n        '#2ca02c',  # Green\n        '#d62728',  # Red\n        '#9467bd',  # Purple\n        '#8c564b',  # Brown\n        '#e377c2',  # Pink\n        '#7f7f7f',  # Gray\n        '#bcbd22',  # Olive\n        '#17becf'   # Cyan\n    ]\n    \n    # (8) Create a mapping of motor ID to color for consistency across slices\n    # This crucial step ensures the same motor gets the same color in each slice,\n    # making it easier to track structures across the Z-axis\n    all_motors = motors_in_tomo\n    unique_motors = {}  # Dictionary to map motor IDs to consistent colors\n    color_counter = 0   # Counter to cycle through the color palette\n    \n    # (8a) Find min and max slice for this tomogram's motors\n    # This helps establish the range for our motor tracking across slices\n    min_slice = min(all_motor_slices)\n    max_slice = max(all_motor_slices)\n    \n    # (8b) Group motors that are close to each other across slices\n    # We scan a range that includes a 5-slice buffer on either side of our actual data\n    # This helps track motors that might appear in slices we're not directly visualizing\n    for z in range(min_slice - 5, max_slice + 6):  \n        # Find all motors in the current slice\n        slice_motors = all_motors[all_motors['Motor axis 0'] == z]\n        \n        # For each motor in this slice, create a unique identifier\n        for _, motor in slice_motors.iterrows():\n            # Create a unique ID based on y and x coordinates\n            # Motors at similar positions in adjacent slices are likely the same structure\n            motor_id = f\"{int(motor['Motor axis 1'])}_{int(motor['Motor axis 2'])}\"\n            \n            # If this is the first time we've seen this motor, assign it a color\n            if motor_id not in unique_motors:\n                unique_motors[motor_id] = motor_colors[color_counter % len(motor_colors)]\n                color_counter += 1\n    \n    # (9) Process and visualize each slice\n    for i, slice_idx in enumerate(slices):\n        # (9a) Get the current subplot from our grid\n        # Calculate row and column indices based on the current slice index\n        row, col = i // cols, i % cols\n        ax = axes[row, col]\n        \n        # (9b) Load the slice image\n        # Construct the file path for the current slice JPEG\n        slice_path = os.path.join(train_dir, tomo_id, f\"slice_{slice_idx:04d}.jpg\")\n        \n        # (9c) Handle missing slice files\n        # If the image file doesn't exist, show an error message instead\n        if not os.path.exists(slice_path):\n            ax.text(0.5, 0.5, f\"Slice {slice_idx} not found\", \n                   ha='center', va='center', fontsize=10)\n            ax.axis('off')\n            continue\n        \n        # (9d) Load and display the image\n        # Read the image file and convert to a numpy array\n        img = np.array(Image.open(slice_path))\n        # Display the grayscale image\n        ax.imshow(img, cmap='gray')\n        \n        # (9e) Filter dataframe for motors in this specific slice\n        # Get only the motors that match both the tomogram ID and the current slice\n        motors = df[(df['tomo_id'] == tomo_id) & (df['Motor axis 0'] == slice_idx)]\n        \n        # (9f) Draw annotation boxes around each motor\n        for j, (_, motor) in enumerate(motors.iterrows()):\n            # Extract the integer pixel coordinates\n            y = int(motor['Motor axis 1'])\n            x = int(motor['Motor axis 2'])\n            # Create a unique ID for this motor based on its position\n            motor_id = f\"{y}_{x}\"\n            \n            # Choose a color for this motor\n            # First try to use a consistent color if this motor has been seen before\n            if motor_id in unique_motors:\n                motor_color = unique_motors[motor_id]\n            else:\n                # Fallback to a sequence-based color if not mapped\n                motor_color = motor_colors[j % len(motor_colors)]\n            \n            # Calculate the box dimensions\n            half_box = box_size // 2\n            \n            # Create a semi-transparent rectangle to highlight the motor\n            rect = Rectangle(\n                (x - half_box, y - half_box),  # Upper-left corner position\n                box_size, box_size,            # Width and height\n                linewidth=2.0,                 # border\n                edgecolor=motor_color,         # Border color\n                facecolor=motor_color,         # Fill with same color as border\n                alpha=0.3                      # Semi-transparent fill for visibility\n            )\n            # Add the rectangle to the plot\n            ax.add_patch(rect)\n            \n            # Add a small label with the motor number\n            #   - For the multi-slice view, we use a compact circular label in the center\n            if show_titlebox:\n                ax.text(\n                    x, y,                          # Center position\n                    f\"{j+1}\",                      # Simple numeric label\n                    color='white',                 # White text for readability\n                    fontsize=7,                    # Small font size to avoid overcrowding\n                    fontweight='bold',             # Bold text for visibility\n                    ha='center',                   # Center-align horizontally\n                    va='center',                   # Center-align vertically\n                    bbox=dict(                     # Background box for contrast\n                        facecolor=motor_color,     # Use the same motor color\n                        alpha=0.8,                 # Mostly opaque for readability\n                        boxstyle='circle',         # Circular label shape\n                        pad=0.1                    # Minimal padding to keep compact\n                    ),\n                    zorder=10                      # Ensure text appears above other elements\n                )\n        \n        # (9g) Add slice information to the subplot title\n        ax.set_title(f\"Slice: {slice_idx}, Motors: {len(motors)}\", fontsize=10)\n        # Hide axis for a cleaner visualization\n        ax.axis('off')\n    \n    # (10) Hide any empty subplots\n    # If our grid has more cells than slices, hide the extras\n    for i in range(n_slices, rows * cols):\n        row, col = i // cols, i % cols\n        axes[row, col].axis('off')\n    \n    # (11) Add an overall title for the figure\n    # Show the slice range if multiple slices, otherwise just the single slice\n    slice_range = f\"Slices: {slices[0]}-{slices[-1]}\" if len(slices) > 1 else f\"Slice: {slices[0]}\"\n    plt.suptitle(f\"Tomogram: {tomo_id} - {slice_range}\", fontsize=16, fontweight='bold')\n    \n    # (12) Adjust layout for optimal viewing\n    # Ensure subplots don't overlap\n    plt.tight_layout()\n    # Make room for the overall title at the top\n    plt.subplots_adjust(top=0.98)\n    \n    # (13) Return the figure and axes for potential further customization\n    return fig, axes\n\n\ndef normalize_and_enhance_contrast(img: np.ndarray, clip_percentile: float = 0.5) -> np.ndarray:\n    \"\"\"Normalizes and enhances contrast in tomogram slices for better visualization.\n    \n    Args:\n        img (np.ndarray): Input image as numpy array.\n        clip_percentile (float, optional): Percentile value for contrast clipping (default: 0.5).\n        \n    Returns:\n        np.ndarray: Contrast-enhanced normalized image.\n    \"\"\"\n    # (1) Determine value range for contrast enhancement\n    p_low = np.percentile(img, clip_percentile)\n    p_high = np.percentile(img, 100 - clip_percentile)\n    \n    # (2) Clip the image values to the determined range\n    img_clipped = np.clip(img, p_low, p_high)\n    \n    # (3) Normalize to 0-1 range\n    img_normalized = (img_clipped - p_low) / (p_high - p_low)\n    \n    return img_normalized","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:23:16.049522Z","iopub.execute_input":"2025-03-13T21:23:16.049892Z","iopub.status.idle":"2025-03-13T21:23:17.566994Z","shell.execute_reply.started":"2025-03-13T21:23:16.049865Z","shell.execute_reply":"2025-03-13T21:23:17.565473Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# visualize_slice_with_motors(\"tomo_226cd8\")  # Shows first slice with a motor\n# visualize_slice_with_motors(\"tomo_226cd8\", slice_idx=169)  # Shows specific slice\n# visualize_multiple_slices(\"tomo_226cd8\")  # Shows all slices with motors\n_f1, _ax1 = visualize_multiple_slices(\"tomo_00e463\")\n_f2, _ax2 = visualize_slice_with_motors(\"tomo_00e463\", slice_idx=225)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<br>\n\n<a id=\"preprocessing\"></a>\n\n<h1 style=\"font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: #d27582;\" id=\"preprocessing\">6&nbsp;&nbsp;DATA PREPROCESSING&nbsp;&nbsp;&nbsp;&nbsp;<a style=\"text-decoration: none; color: #6c0f1c;\" href=\"#toc\">&#10514;</a></h1>\n\n<br>\n","metadata":{}},{"cell_type":"code","source":"def prepare_yolo_dataset(\n    train_dir: str,\n    train_labels_path: str,\n    yolo_dataset_dir: str,\n    yolo_images_train: str,\n    yolo_images_val: str,\n    yolo_labels_train: str,\n    yolo_labels_val: str,\n    trust: int = 4,\n    box_size: int = 24,\n    train_split: float = 0.8,\n    random_seed: int = 42,\n    return_labels_df: bool = True\n) -> dict[str, Any]:\n    \"\"\"\n    Prepare the complete YOLO dataset from tomograms.\n    \n    This is the main function that orchestrates the entire dataset preparation process.\n    \n    Args:\n        train_dir (str): Directory containing the raw tomogram training data\n        train_labels_path (str): Path to the CSV file with motor labels\n        yolo_dataset_dir (str): Base directory for the YOLO dataset\n        yolo_images_train (str): Directory for training images\n        yolo_images_val (str): Directory for validation images\n        yolo_labels_train (str): Directory for training labels\n        yolo_labels_val (str): Directory for validation labels\n        trust (int, optional): Number of slices above and below center slice to include\n        box_size (int, optional): Size of bounding box in pixels for annotations\n        train_split (float, optional): Fraction of data to use for training (0.0-1.0)\n        random_seed (int, optional): Random seed for reproducibility\n        return_labels_df (bool, optional): Whether to return the loaded labels dataframe or not.\n        \n    Returns:\n        dict[str, Any]:\n            Summary statistics and paths\n        \n    Raises:\n        Various exceptions if processing fails\n    \"\"\"\n    # (1) Validate parameters\n    if not 0.0 < train_split < 1.0:\n        raise ValueError(\"train_split must be between 0.0 and 1.0\")\n    \n    if trust < 0:\n        raise ValueError(\"trust must be a non-negative integer\")\n        \n    if box_size <= 0:\n        raise ValueError(\"box_size must be a positive integer\")\n    \n    # (2) Create necessary directories\n    create_dataset_directories(\n        yolo_images_train, \n        yolo_images_val, \n        yolo_labels_train, \n        yolo_labels_val\n    )\n    \n    # (3) Load and validate the labels\n    labels_df = validate_labels_file(train_labels_path)\n    \n    # (4) Count total number of motors for reporting\n    total_motors = labels_df['Number of motors'].sum()\n    print(f\"\\nTOTAL NUMBER OF MOTORS IN THE DATASET: {total_motors}\")\n    \n    # (5) Split tomograms into training and validation sets\n    train_tomos, val_tomos = split_tomograms(\n        labels_df, \n        train_split=train_split, \n        random_seed=random_seed\n    )\n    \n    # (6) Process training tomograms\n    train_slices, train_motors = process_tomogram_set(\n        labels_df,\n        train_tomos, \n        train_dir,\n        yolo_images_train, \n        yolo_labels_train, \n        \"training\",\n        trust=trust,\n        box_size=box_size\n    )\n    \n    # (7) Process validation tomograms\n    val_slices, val_motors = process_tomogram_set(\n        labels_df,\n        val_tomos, \n        train_dir,\n        yolo_images_val, \n        yolo_labels_val, \n        \"validation\",\n        trust=trust,\n        box_size=box_size\n    )\n    \n    # (8) Create YAML configuration file for YOLO\n    yaml_path = create_yaml_config(yolo_dataset_dir)\n    \n    # (9) Create and populate summary statistics\n    stats = {\n        \"train_tomograms\": len(train_tomos),\n        \"val_tomograms\": len(val_tomos),\n        \"train_motors\": train_motors,\n        \"val_motors\": val_motors,\n        \"train_slices\": train_slices,\n        \"val_slices\": val_slices,\n        \"dataset_dir\": yolo_dataset_dir,\n        \"yaml_path\": yaml_path\n    }\n    \n    # (10) Print summary information\n    print(f\"\\nProcessing Summary:\")\n    print(f\"- Train set: {stats['train_tomograms']} tomograms, {train_motors} motors, {train_slices} slices\")\n    print(f\"- Validation set: {stats['val_tomograms']} tomograms, {val_motors} motors, {val_slices} slices\")\n    print(f\"- Total: {stats['train_tomograms'] + stats['val_tomograms']} tomograms, \"\n          f\"{train_motors + val_motors} motors, {train_slices + val_slices} slices\")\n    \n    # (11) Return statistics dictionary\n    if not return_labels_df:\n        return stats\n\n    # (12) Optionally return the statistics dictionary and the labels dataframe\n    return stats, labels_df\n\n# Run the preprocessing\nsummary, labels_df = prepare_yolo_dataset(\n    train_dir=TRAIN_DIR,\n    train_labels_path=TRAIN_LABELS_PATH,\n    yolo_dataset_dir=YOLO_DATASET_DIR,\n    yolo_images_train=YOLO_IMAGES_TRAIN,\n    yolo_images_val=YOLO_IMAGES_VAL,\n    yolo_labels_train=YOLO_LABELS_TRAIN,\n    yolo_labels_val=YOLO_LABELS_VAL,\n    trust=TRUST,\n    box_size=BOX_SIZE,\n    train_split=TRAIN_SPLIT,\n    random_seed=42\n)\n\nprint(f\"\\nPreprocessing Complete:\")\nprint(f\"  - Training data: {summary['train_tomograms']} tomograms, {summary['train_motors']} motors, {summary['train_slices']} slices\")\nprint(f\"  - Validation data: {summary['val_tomograms']} tomograms, {summary['val_motors']} motors, {summary['val_slices']} slices\")\nprint(f\"  - Dataset directory: {summary['dataset_dir']}\")\nprint(f\"  - YAML configuration: {summary['yaml_path']}\")\nprint(f\"\\nReady for YOLO training!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T02:57:50.141188Z","iopub.execute_input":"2025-03-12T02:57:50.141538Z","execution_failed":"2025-03-12T02:58:05.102Z"}},"outputs":[],"execution_count":null}]}