{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install streamlit","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:40:50.769221Z","iopub.execute_input":"2026-02-22T14:40:50.770077Z","iopub.status.idle":"2026-02-22T14:40:56.938377Z","shell.execute_reply.started":"2026-02-22T14:40:50.770047Z","shell.execute_reply":"2026-02-22T14:40:56.937665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pyngrok --quiet\n!pip install reportlab --quiet\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:40:56.939668Z","iopub.execute_input":"2026-02-22T14:40:56.939908Z","iopub.status.idle":"2026-02-22T14:41:04.509021Z","shell.execute_reply.started":"2026-02-22T14:40:56.939888Z","shell.execute_reply":"2026-02-22T14:41:04.508097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/app.py\n# ================================================================\n# 🧬 Phase 1 – Flagellar Motor 3D Analysis (Research-Ready Baseline)\n# ================================================================\n\nimport os, gc, warnings, random, time\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nfrom scipy import ndimage\nfrom sklearn.cluster import DBSCAN\nfrom sklearn.neighbors import NearestNeighbors\nfrom PIL import Image as PILImage\n\nimport torch\nimport torch.nn as nn\n\nimport streamlit as st\nimport plotly.graph_objects as go\nimport plotly.express as px\n\nfrom typing import Tuple\n\n# ---------------------------------------------------------------\n# ✅ 1. Deterministic Reproducibility\n# ---------------------------------------------------------------\ndef set_seed(seed: int = 42):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False  # enables mixed-data speeds\n    torch.backends.cudnn.benchmark = True\n\nset_seed(42)\n\n# ---------------------------------------------------------------\n# ✅ 2. Optional Import: skimage (surface reconstruction)\n# ---------------------------------------------------------------\ntry:\n    from skimage import measure\nexcept Exception:\n    measure = None\n    warnings.warn(\"skimage.measure unavailable — 3D marching_cubes disabled.\", UserWarning)\n\n# ---------------------------------------------------------------\n# ✅ 3. Global Configuration\n# ---------------------------------------------------------------\nDATA_ROOT = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"\nTRAIN_FOLDER = \"train\"\nGT_CSV = \"train_labels.csv\"\nCACHE_DIR = \"/kaggle/working/cache_phase1\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n\n# --- Compute / Visualization Control ---\nMAX_POINTS, MAX_VIS_POINTS = 15000, 100000\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# --- Training Defaults ---\nPATCH_SIZE = 64\nPATCH_STRIDE = 32\nTRAIN_PATCHES_PER_EPOCH = 64\nBATCH_SIZE = 1\nEPOCHS = 5\nLEARNING_RATE = 1e-3\nSIGMA_HEATMAP = 1.2\nTHRESH_STD = 0.4\nDBSCAN_EPS = 2.0\n\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.enabled = True\n\n# ---------------------------------------------------------------\n# ✅ 4. Streamlit Environment Setup\n# ---------------------------------------------------------------\nst.set_page_config(page_title=\"Flagellar Motor 3D Analysis\", layout=\"wide\")\nst.title(\"🧬 Phase 1 – Flagellar Motor 3D Analysis (Research-Ready Baseline)\")\n\n# ---------------------------------------------------------------\n# ✅ 5. Core I/O Utilities (robust + supports .npy tomograms)\n# ---------------------------------------------------------------\ndef find_tomograms(root: str | Path):\n    \"\"\"\n    Return tomogram folders or volume files under dataset path.\n    Supports both subfolders (with .jpg/.png) and direct .npy volumes.\n    \"\"\"\n    root = Path(root)\n    p = root / TRAIN_FOLDER if (root / TRAIN_FOLDER).exists() else root\n    if not p.exists():\n        st.warning(f\"⚠️ Dataset path not found: {p}\")\n        return []\n    tomos = []\n    for d in p.iterdir():\n        if d.is_dir():\n            imgs = list(d.glob(\"*.jpg\")) + list(d.glob(\"*.png\"))\n            if len(imgs) > 0:\n                tomos.append(d)\n        elif d.suffix == \".npy\":\n            tomos.append(d)\n    if not tomos:\n        st.warning(f\"⚠️ No tomograms found under {p}\")\n    return tomos\n\n\ndef save_memmap_and_return(path: Path, arr: np.ndarray) -> np.memmap:\n    \"\"\"Save NumPy array and reopen as memmap for large-volume streaming.\"\"\"\n    np.save(path, arr)\n    return np.load(path, mmap_mode=\"r\")\n\n\n@st.cache_data(show_spinner=False)\ndef load_volume_from_jpegs_cached(tomo_path: str | Path):\n    \"\"\"Load stack of 2D slices or .npy volume into a memmap.\"\"\"\n    tomo_p = Path(tomo_path)\n    cache_file = Path(CACHE_DIR) / f\"{tomo_p.stem}_raw.npy\"\n\n    # ✅ Case 1: cached memmap\n    if cache_file.exists():\n        try:\n            return np.load(cache_file, mmap_mode=\"r\")\n        except Exception:\n            cache_file.unlink(missing_ok=True)\n\n    # ✅ Case 2: directory of slices\n    if tomo_p.is_dir():\n        slices = sorted(tomo_p.glob(\"*.jpg\")) + sorted(tomo_p.glob(\"*.png\"))\n        if not slices:\n            st.error(f\"❌ No image slices found in {tomo_p}\")\n            return None\n        vol = np.stack(\n            [np.array(PILImage.open(s)) for s in slices],\n            axis=0\n        ).astype(np.float32)\n\n        return save_memmap_and_return(cache_file, vol)\n\n    # ✅ Case 3: direct .npy file\n    if tomo_p.is_file() and tomo_p.suffix == \".npy\":\n        return np.load(tomo_p, mmap_mode=\"r\")\n\n    st.error(f\"❌ Unsupported tomogram path: {tomo_p}\")\n    return None\n\n# ---------------------------------------------------------------\n# ✅ 5b. Cached Preprocessing (ALWAYS disk-backed)\n# ---------------------------------------------------------------\n@st.cache_data(show_spinner=False)\ndef preprocess_volume_cached_memmap(volume_memmap):\n    \"\"\"\n    Ensure volume is disk-backed memmap with valid filename.\n    This guarantees downstream compatibility.\n    \"\"\"\n    if volume_memmap is None:\n        return None\n\n    cache_path = Path(CACHE_DIR) / \"preprocessed_volume.npy\"\n\n    # Case 1: already a memmap with filename\n    if isinstance(volume_memmap, np.memmap) and volume_memmap.filename is not None:\n        return volume_memmap\n\n    # Case 2: memmap without filename OR ndarray → force save\n    arr = np.array(volume_memmap, dtype=np.float32)\n    np.save(cache_path, arr)\n\n    return np.load(cache_path, mmap_mode=\"r\")\n\ndef ensure_disk_memmap(mm):\n    if isinstance(mm, np.memmap) and mm.filename is not None:\n        return mm\n    path = Path(CACHE_DIR) / \"forced_disk_memmap.npy\"\n    np.save(path, np.array(mm, dtype=np.float32))\n    return np.load(path, mmap_mode=\"r\")\n\n\n\n# ---------------------------------------------------------------\n# ✅ 6a. Tomogram Selection + Volume Preparation (CRITICAL)\n# ---------------------------------------------------------------\ntomos = find_tomograms(DATA_ROOT)\nif not tomos:\n    st.stop()\n\nselected_tomo = st.selectbox(\n    \"Select Tomogram\",\n    tomos,\n    format_func=lambda p: p.name if isinstance(p, Path) else str(p)\n)\n\n# ---- Load selected tomogram ----\nraw_memmap = load_volume_from_jpegs_cached(selected_tomo)\nif raw_memmap is None:\n    st.error(\"❌ Failed to load tomogram.\")\n    st.stop()\n\n# ---- Preprocess (cached, disk-backed) ----\nvol_memmap = preprocess_volume_cached_memmap(raw_memmap)\nif vol_memmap is None:\n    st.error(\"❌ Preprocessing failed.\")\n    st.stop()\n\n# ---- Final safety: enforce disk-backed memmap ----\nvol_memmap = ensure_disk_memmap(vol_memmap)\n\n# ---- HARD GUARD: filename must exist ----\nif not isinstance(vol_memmap, np.memmap) or vol_memmap.filename is None:\n    st.error(\"❌ Phase 1 volume is not disk-backed.\")\n    st.stop()\n\n# ---- Safe to use filename now ----\nproc_cache_path = Path(vol_memmap.filename)\nst.success(f\"✅ Volume ready: {vol_memmap.shape}\")\n\n\n\n# ---------------------------------------------------------------\n# ✅ 6. Lightweight Geometry Utilities\n# ---------------------------------------------------------------\ndef fallback_pointcloud(volume_memmap, downsample=4, threshold=None) -> np.ndarray:\n    \"\"\"Quick voxel-threshold to approximate 3D coordinates (fallback).\"\"\"\n    if volume_memmap is None:\n        return np.zeros((0, 3))\n    mm = np.load(volume_memmap, mmap_mode=\"r\") if isinstance(volume_memmap, (str, Path)) else volume_memmap\n    if threshold is None:\n        sample = []\n        step = max(1, mm.shape[0] // 10)\n        for zi in range(0, mm.shape[0], step):\n            sample.append(mm[zi])\n        sample = np.concatenate([s.ravel() for s in sample])\n        threshold = sample.mean() + 0.5 * sample.std()\n        del sample\n    coords_list = []\n    for zi in range(mm.shape[0]):\n        sl = mm[zi]\n        ys, xs = np.nonzero(sl > threshold)\n        if ys.size > 0:\n            zcol = np.full_like(ys, zi, dtype=float)\n            coords_list.append(np.stack([zcol, ys.astype(float), xs.astype(float)], axis=1))\n    if not coords_list:\n        return np.zeros((0, 3))\n    coords = np.vstack(coords_list)\n    if coords.shape[0] > MAX_POINTS:\n        idx = np.random.choice(coords.shape[0], size=MAX_POINTS, replace=False)\n        coords = coords[idx]\n    return coords[:, [2, 1, 0]].astype(float)\n\n@st.cache_data(show_spinner=False)\ndef marching_cubes_mesh_adaptive(volume_memmap, level: float = None, target_max_vertices: int = MAX_POINTS):\n    \"\"\"Adaptive marching_cubes mesh with pre-emptive downsampling to prevent OOM.\"\"\"\n    if measure is None:\n        return None, None\n        \n    mm = np.load(volume_memmap, mmap_mode=\"r\") if isinstance(volume_memmap, (str, Path)) else volume_memmap\n    if mm is None or mm.size == 0:\n        return None, None\n        \n    # ✅ FIX: Pre-emptive downsampling to prevent RAM blowout!\n    # Cap processing at ~5 million voxels for safety on Kaggle\n    target_voxels = 5_000_000 \n    factor = int(np.ceil((mm.size / target_voxels) ** (1/3)))\n    factor = max(1, factor)\n    \n    # Efficient strided read from disk (memmap) into RAM\n    vol = np.array(mm[::factor, ::factor, ::factor], dtype=np.float32)\n    \n    if level is None:\n        level = float(vol.mean() + 0.5 * vol.std())\n        \n    try:\n        verts, faces, _, _ = measure.marching_cubes(vol, level=level)\n        verts *= factor  # Scale coordinates back up to original volume space\n    except Exception as e:\n        warnings.warn(f\"marching_cubes failed even after downsampling: {e}\")\n        return None, None\n        \n    # Safety cap for Plotly rendering speed in browser\n    if verts.shape[0] > target_max_vertices:\n        # If we subsample vertices, faces become invalid, so we drop faces and just return a point cloud\n        faces = np.zeros((0, 3), dtype=int)\n        idx = np.random.choice(verts.shape[0], size=target_max_vertices, replace=False)\n        verts = verts[idx]\n        \n    return verts[:, [2, 1, 0]], faces\n\ndef cluster_points(coords: np.ndarray, eps=DBSCAN_EPS):\n    \"\"\"DBSCAN clustering (GPU-safe).\"\"\"\n    if coords is None or coords.shape[0] == 0:\n        return np.zeros((0, 3))\n    try:\n        clustering = DBSCAN(eps=eps, min_samples=2).fit(coords)\n        return coords[clustering.labels_ >= 0]\n    except Exception:\n        return coords\n\ndef match_predictions_to_gt(pred_pts, gt_pts, radius=5.0):\n    \"\"\"Simple nearest-neighbor matching (used by Part 2 metrics).\"\"\"\n    if len(pred_pts) == 0 or len(gt_pts) == 0:\n        return 0, len(gt_pts), np.array([])\n    nbrs = NearestNeighbors(n_neighbors=1).fit(pred_pts)\n    dists, _ = nbrs.kneighbors(gt_pts)\n    matched = dists[:, 0] <= radius\n    return int(matched.sum()), int(len(gt_pts)), dists[:, 0]\n\n# ---------------------------------------------------------------\n# ✅ 6b. Physical-Unit Coordinate Normalization (nm-scale)\n# ---------------------------------------------------------------\nfrom typing import Tuple\n\ndef get_default_voxel_size(tomo_name: str | Path) -> Tuple[float, float, float]:\n    \"\"\"\n    Returns voxel size (z_nm, y_nm, x_nm) for a tomogram.\n    Replace with metadata lookup if available.\n    Example defaults below assume anisotropy typical in cryo-ET (z≈4–5× coarser).\n    \"\"\"\n    name = str(tomo_name).lower()\n    # ⚙️ Customize per dataset / MotorBench calibration\n    if \"byu\" in name or \"motor\" in name:\n        return (5.0, 1.0, 1.0)  # 5 nm along Z, 1 nm along Y/X\n    return (1.0, 1.0, 1.0)  # fallback isotropic\n\ndef vox_to_nm(coords: np.ndarray, voxel_size_nm: Tuple[float, float, float]):\n    \"\"\"Convert voxel indices (z,y,x) → nanometers.\"\"\"\n    if coords is None or coords.size == 0:\n        return coords\n    z_nm, y_nm, x_nm = voxel_size_nm\n    # ✅ FIX: Match the scale array to the Z, Y, X coordinate format\n    scale = np.array([z_nm, y_nm, x_nm], dtype=float)\n    return coords * scale\n\ndef nm_to_vox(coords_nm: np.ndarray, voxel_size_nm: Tuple[float, float, float]):\n    \"\"\"Convert nanometer coordinates back to voxel units.\"\"\"\n    if coords_nm is None or coords_nm.size == 0:\n        return coords_nm\n    z_nm, y_nm, x_nm = voxel_size_nm\n    # ✅ FIX: Match the scale array to the Z, Y, X coordinate format\n    inv = np.array([1/z_nm, 1/y_nm, 1/x_nm], dtype=float)\n    return coords_nm * inv\n\ndef compute_metrics_nm(pred_pts: np.ndarray, gt_pts: np.ndarray,\n                       voxel_size_nm: Tuple[float, float, float], radius_nm: float = 10.0):\n    \"\"\"\n    Compute metrics (precision, recall, F2, mean_dist_nm) in physical nanometer units.\n    \"\"\"\n    if len(gt_pts) == 0:\n        return dict(precision=0, recall=0, f2=0, mean_dist_nm=0, matched=0)\n    if len(pred_pts) == 0:\n        return dict(precision=0, recall=0, f2=0, mean_dist_nm=np.nan, matched=0)\n\n    from sklearn.neighbors import NearestNeighbors\n    pred_nm = vox_to_nm(pred_pts, voxel_size_nm)\n    gt_nm   = vox_to_nm(gt_pts, voxel_size_nm)\n    nbrs = NearestNeighbors(n_neighbors=1).fit(pred_nm)\n    dists, _ = nbrs.kneighbors(gt_nm)\n    matched = dists[:, 0] <= radius_nm\n\n    tp = matched.sum()\n    fp = len(pred_nm) - tp\n    fn = len(gt_nm) - tp\n    precision = tp / (tp + fp + 1e-8)\n    recall    = tp / (tp + fn + 1e-8)\n    f2        = (1 + 2**2) * (precision * recall) / (4 * precision + recall + 1e-8)\n\n    return dict(\n        precision=float(precision),\n        recall=float(recall),\n        f2=float(f2),\n        mean_dist_nm=float(dists.mean()),\n        matched=int(tp)\n    )\n\n\n# ---------------------------------------------------------------\n# ✅ 7. Visualization (shared by later parts)\n# ---------------------------------------------------------------\ndef safe_sample_points(arr: np.ndarray, max_points: int = MAX_VIS_POINTS):\n    if arr is None or arr.size == 0:\n        return np.zeros((0, 3))\n    if arr.shape[0] > max_points:\n        idx = np.random.choice(arr.shape[0], size=max_points, replace=False)\n        return arr[idx]\n    return arr\n\ndef plot_volume_mesh_and_points(mesh_verts, mesh_faces,\n                                points=None, gt_pts=None, pred_pts=None,\n                                title=\"3D Mesh + Points\"):\n    \"\"\"3D plot combining mesh and predicted points.\"\"\"\n    points_vis = safe_sample_points(points)\n    gt_vis = safe_sample_points(gt_pts)\n    mesh_vis = safe_sample_points(mesh_verts, max_points=MAX_POINTS) if mesh_verts is not None else None\n    fig = go.Figure()\n    if mesh_vis is not None and mesh_faces is not None and len(mesh_faces) > 0:\n        i, j, k = mesh_faces.T\n        fig.add_trace(go.Mesh3d(x=mesh_vis[:, 0], y=mesh_vis[:, 1], z=mesh_vis[:, 2],\n                                i=i, j=j, k=k, opacity=0.2, color=\"lightblue\", name=\"volume mesh\"))\n    for pts, name, color, size in [(points_vis, \"volume points\", \"blue\", 1),\n                                   (gt_vis, \"ground-truth\", \"green\", 4),\n                                   (pred_pts, \"predicted\", \"red\", 4)]:\n        if pts is not None and getattr(pts, \"shape\", (0,))[0] > 0:\n            fig.add_trace(go.Scatter3d(x=pts[:, 0], y=pts[:, 1], z=pts[:, 2],\n                                       mode=\"markers\", marker=dict(size=size, color=color), name=name))\n    fig.update_layout(scene=dict(aspectmode=\"data\"), title=title, width=900, height=700)\n    return fig\n\ndef plot_slice_with_points_memmap(memmap_path: Path, z_index: int,\n                                  points=None, gt_pts=None, pred_pts=None,\n                                  title=\"Slice View\"):\n    \"\"\"Interactive 2D slice viewer with overlayed points.\"\"\"\n    mm = np.load(memmap_path, mmap_mode=\"r\")\n    z_index = int(max(0, min(z_index, mm.shape[0] - 1)))\n    slice_img = np.array(mm[z_index, :, :])\n\n    def pts_on_slice(pts):\n        if pts is None or pts.size == 0:\n            return np.zeros((0, 3))\n        z_inds = np.round(pts[:, 2]).astype(int)\n        return pts[z_inds == z_index]\n\n    pts = pts_on_slice(points)\n    gt = pts_on_slice(gt_pts)\n    pred = pts_on_slice(pred_pts)\n\n    fig = px.imshow(slice_img, color_continuous_scale=\"gray\", origin=\"lower\")\n    if pts.shape[0] > 0:\n        fig.add_scatter(x=pts[:, 0], y=pts[:, 1], mode=\"markers\",\n                        marker=dict(size=3, color=\"blue\"), name=\"volume\")\n    if gt.shape[0] > 0:\n        fig.add_scatter(x=gt[:, 0], y=gt[:, 1], mode=\"markers\",\n                        marker=dict(size=5, color=\"green\"), name=\"GT\")\n    if pred.shape[0] > 0:\n        fig.add_scatter(x=pred[:, 0], y=pred[:, 1], mode=\"markers\",\n                        marker=dict(size=5, color=\"red\"), name=\"Pred\")\n    fig.update_layout(title=f\"{title} – Z={z_index}\", width=700, height=700)\n    return fig\n\n# ---------------------------------------------------------------\n# 🔒 GLOBAL VOLUME GUARD — prevents Streamlit rerun crashes\n# ---------------------------------------------------------------\nif (\n    \"vol_memmap\" not in globals()\n    or vol_memmap is None\n    or not isinstance(vol_memmap, np.memmap)\n    or vol_memmap.filename is None\n):\n    st.warning(\"⚠️ No valid Phase 1 volume — downstream stages disabled.\")\n    st.stop()\n\n\n# ================================================================\n# Part 2 – Tiny3DCNN + Confidence Filtering & Extended Training\n# ================================================================\n\nimport os, gc, time, random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom pathlib import Path\nimport pandas as pd\nimport streamlit as st\nimport plotly.express as px\nfrom sklearn.neighbors import NearestNeighbors\nfrom scipy import ndimage\n\n# ---------------- Constants ----------------\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nPATCH_SIZE = 64\nPATCH_STRIDE = 32\nTRAIN_PATCHES_PER_EPOCH = 64\nEPOCHS = 10   # ⚗️ extended training for smoother heatmaps\nLEARNING_RATE = 1e-3\nSIGMA_HEATMAP = 1.2\nTHRESH_STD = 0.4\nDBSCAN_EPS = 2.0\nMAX_POINTS = 15000\nCACHE_DIR = \"/kaggle/working/cache_phase1\"\nLOG_CSV = Path(CACHE_DIR) / \"benchmark_runs.csv\"\n\n# ---------------- Seeding ----------------\ndef set_seed(seed=42):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\nset_seed(42)\n\n# ---------------- Model ----------------\nclass Tiny3DCNN(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv3d(1, 16, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(16, 32, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(32, 16, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(16, 1, 1), nn.Sigmoid()\n        )\n    def forward(self, x): return self.net(x)\n\n# ---------------- Metrics ----------------\ndef compute_metrics(pred_pts, gt_pts, radius=5.0):\n    if len(gt_pts) == 0:\n        return dict(precision=0, recall=0, f2=0, mean_dist=0, matched=0)\n    if len(pred_pts) == 0:\n        return dict(precision=0, recall=0, f2=0, mean_dist=np.nan, matched=0)\n    nbrs = NearestNeighbors(n_neighbors=1).fit(pred_pts)\n    dists, _ = nbrs.kneighbors(gt_pts)\n    matched = dists[:, 0] <= radius\n    tp = matched.sum(); fp = len(pred_pts) - tp; fn = len(gt_pts) - tp\n    precision = tp / (tp + fp + 1e-8)\n    recall    = tp / (tp + fn + 1e-8)\n    f2        = (1 + 2**2) * (precision * recall) / (4 * precision + recall + 1e-8)\n    return dict(precision=float(precision), recall=float(recall),\n                f2=float(f2), mean_dist=float(dists.mean()), matched=int(tp))\n\n# ---------------- Patches ----------------\ndef extract_random_patches(volume_arr, target, n_patches, patch_size):\n    \"\"\"Extract patches, ensuring 50% of them contain ground truth targets if available.\"\"\"\n    Z, Y, X = volume_arr.shape\n    patches_v, patches_t = [], []\n    \n    # 🧠 FIX: Use a relative threshold based on the actual Gaussian peak\n    t_max = target.max()\n    has_positives = False\n    if t_max > 1e-5:\n        # Grab voxels that are at least 10% of the maximum peak\n        pos_z, pos_y, pos_x = np.nonzero(target > (t_max * 0.1))\n        has_positives = len(pos_z) > 0\n\n    for i in range(n_patches):\n        # 50% chance to aggressively sample around a known target\n        if has_positives and (i % 2 == 0):\n            idx = np.random.randint(0, len(pos_z))\n            zi = max(0, min(pos_z[idx] - patch_size // 2, Z - patch_size))\n            yi = max(0, min(pos_y[idx] - patch_size // 2, Y - patch_size))\n            xi = max(0, min(pos_x[idx] - patch_size // 2, X - patch_size))\n        else:\n            # Random background sampling\n            zi = np.random.randint(0, max(1, Z - patch_size + 1))\n            yi = np.random.randint(0, max(1, Y - patch_size + 1))\n            xi = np.random.randint(0, max(1, X - patch_size + 1))\n            \n        patches_v.append(volume_arr[zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size])\n        patches_t.append(target[zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size])\n        \n    return np.stack(patches_v), np.stack(patches_t)\n\n# ---------------- Inference (with patch confidence) ----------------\ndef sliding_window_inference(volume_arr, model, patch_size=PATCH_SIZE, stride=PATCH_STRIDE):\n    model.eval()\n    Z, Y, X = volume_arr.shape\n    accum, norm, conf_map = (np.zeros((Z, Y, X), np.float32) for _ in range(3))\n    \n    with torch.no_grad():\n        for zi in range(0, Z - patch_size + 1, stride):\n            for yi in range(0, Y - patch_size + 1, stride):\n                for xi in range(0, X - patch_size + 1, stride):\n                    patch = volume_arr[zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size].astype(np.float32)\n                    inp = torch.from_numpy(patch[None, None, ...]).to(DEVICE)\n                    \n                    with torch.amp.autocast('cuda', enabled=(DEVICE == \"cuda\")):\n                        out = model(inp)\n                        \n                    out_np = out[0, 0].cpu().numpy()\n                    \n                    # 🧠 FIX: Use .max() instead of .mean() for sparse 3D keypoints!\n                    # We only care if there is a strong peak *somewhere* in the patch.\n                    conf = float(out_np.max())\n                    if conf < 0.1:  # skip patches that are completely empty\n                        continue\n                        \n                    accum[zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size] += out_np\n                    norm [zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size] += 1.0\n                    conf_map[zi:zi+patch_size, yi:yi+patch_size, xi:xi+patch_size] += conf\n                    \n                    del inp, out, out_np\n                    torch.cuda.empty_cache()\n                    \n    norm[norm == 0] = 1.0\n    conf_map /= norm\n    import gc\n    gc.collect(); torch.cuda.empty_cache()\n    return accum / norm, conf_map\n\n# ---------------- Train + Infer + Log ----------------\n\ndef run_tiny_cnn(volume_memmap_or_arr, gt_points, epochs=EPOCHS, preset=\"balanced\",\n                 stride_factor=1.0):\n    \"\"\"\n    Train Tiny3DCNN, infer on tomogram, and compute nm-scale benchmark metrics.\n    Includes patch-wise confidence filtering for better recall balance.\n    \"\"\"\n    t0_total = time.time()\n    gpu_start_mem = torch.cuda.memory_allocated(DEVICE) if DEVICE == \"cuda\" else 0\n\n    # ---- Load volume ----\n    if isinstance(volume_memmap_or_arr, (str, Path)):\n        mm = np.load(volume_memmap_or_arr, mmap_mode='r')\n        vol = np.array(mm)\n    else:\n        vol = np.array(volume_memmap_or_arr)\n    if vol.size == 0:\n        st.warning(\"Empty volume — cannot run CNN.\")\n        return np.zeros((0,3)), {}\n\n    # ---- Build target heatmap ----\n    Z, Y, X = vol.shape\n    target = np.zeros_like(vol, dtype=np.float32)\n    for z, y, x in np.asarray(gt_points).astype(int):\n        if 0 <= z < Z and 0 <= y < Y and 0 <= x < X:\n            target[z, y, x] = 1.0\n    target = ndimage.gaussian_filter(target, sigma=SIGMA_HEATMAP)\n\n    # ---- Model setup ----\n    model = Tiny3DCNN().to(DEVICE)\n    optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n    loss_fn = nn.MSELoss()\n    scaler = torch.cuda.amp.GradScaler(enabled=(DEVICE == \"cuda\"))\n    total_steps = epochs * TRAIN_PATCHES_PER_EPOCH\n    pb = st.progress(0)\n    losses = []\n\n    # ---- Training ----\n    for ep in range(epochs):\n        patches_v, patches_t = extract_random_patches(vol, target, TRAIN_PATCHES_PER_EPOCH, PATCH_SIZE)\n        model.train()\n        for i in range(patches_v.shape[0]):\n            inp = torch.from_numpy(patches_v[i:i+1]).unsqueeze(1).to(DEVICE)\n            tgt = torch.from_numpy(patches_t[i:i+1]).unsqueeze(1).to(DEVICE)\n            optimizer.zero_grad()\n            with torch.cuda.amp.autocast(enabled=(DEVICE == \"cuda\")):\n                out = model(inp)\n                loss = loss_fn(out, tgt)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            losses.append(loss.item())\n            pb.progress(int(100 * ((ep * TRAIN_PATCHES_PER_EPOCH + i + 1) / total_steps)))\n            del inp, tgt, out, loss\n        gc.collect(); torch.cuda.empty_cache()\n\n    # ---- Inference + confidence ----\n    st.info(\"Running inference with patch confidence filtering...\")\n    recon, conf_map = sliding_window_inference(\n        vol, model, patch_size=PATCH_SIZE, stride=int(PATCH_STRIDE * stride_factor)\n    )\n\n    # ---- Adaptive thresholding (confidence-weighted) ----\n    dynamic_thresh = (recon.mean() + THRESH_STD * recon.std()) * (1 - 0.15 * conf_map.mean())\n    coords = np.array(np.nonzero(recon > dynamic_thresh)).T\n    if coords.shape[0] > MAX_POINTS:\n        coords = coords[np.random.choice(coords.shape[0], MAX_POINTS, replace=False)]\n    coords_xyz = coords[:, [2, 1, 0]].astype(float)\n    clustered = cluster_points(coords_xyz, eps=DBSCAN_EPS)\n\n    # ---- Metrics (nm-normalized) ----\n    voxel_size_nm = get_default_voxel_size(\"byu_motor\")  # can adapt per dataset\n    metrics = compute_metrics_nm(clustered, gt_points, voxel_size_nm, radius_nm=10.0)\n\n    metrics.update({\n        \"train_time_s\": round(time.time() - t0_total, 2),\n        \"gpu_mem_MB\": round((torch.cuda.max_memory_allocated(DEVICE) - gpu_start_mem) / (1024**2), 2)\n                      if DEVICE == \"cuda\" else 0,\n        \"mean_train_loss\": float(np.mean(losses)),\n        \"pred_count\": len(clustered),\n        \"gt_count\": len(gt_points),\n        \"confidence_mean\": float(conf_map.mean()),\n        \"voxel_size_nm_z\": voxel_size_nm[0],\n        \"voxel_size_nm_y\": voxel_size_nm[1],\n        \"voxel_size_nm_x\": voxel_size_nm[2],\n        \"mean_dist_nm\": metrics.get(\"mean_dist_nm\", np.nan)\n    })\n\n    # ---- Save ----\n    out_path = Path(CACHE_DIR) / f\"pred_{int(np.random.rand()*1e9)}.npz\"\n    np.savez_compressed(out_path, points=clustered)\n    metrics[\"saved_path\"] = str(out_path)\n\n    # ---- Log ----\n    df_log = pd.DataFrame([metrics])\n    if LOG_CSV.exists():\n        prev = pd.read_csv(LOG_CSV)\n        df_log = pd.concat([prev, df_log], ignore_index=True)\n    df_log.to_csv(LOG_CSV, index=False)\n\n    # ---- Report ----\n    st.success(\n        f\"✅ Done — F₂={metrics['f2']:.3f}, \"\n        f\"Precision={metrics['precision']:.2f}, Recall={metrics['recall']:.2f}, \"\n        f\"Mean Dist ≈ {metrics['mean_dist_nm']:.2f} nm\"\n    )\n    st.write(\"Mean Patch Confidence:\", f\"{metrics['confidence_mean']:.3f}\")\n    st.dataframe(pd.DataFrame([metrics]))\n    return clustered, metrics\n\n# ================================================================\n# 📊 ΔF₂ / Precision / Recall Visual Analytics (Preset + Model-aware)\n# ================================================================\n\nimport plotly.express as px\n\nst.subheader(\"📈 ΔF₂ / Precision / Recall Analysis Across Runs\")\n\nlog_path = Path(CACHE_DIR) / \"benchmark_runs.csv\"\nif log_path.exists():\n    df = pd.read_csv(log_path)\n\n    if len(df) >= 2:\n        df = df.sort_values(\"train_time_s\").reset_index(drop=True)\n\n        # Compute percentage deltas\n        df[\"ΔF2_%\"] = df[\"f2\"].pct_change() * 100\n        df[\"ΔRecall_%\"] = df[\"recall\"].pct_change() * 100\n        df[\"ΔPrecision_%\"] = df[\"precision\"].pct_change() * 100\n\n        # Fill in missing columns gracefully\n        if \"preset\" not in df.columns:\n            df[\"preset\"] = \"balanced\"\n        if \"model\" not in df.columns:\n            df[\"model\"] = \"Tiny3DCNN\"\n\n        # 📊 F₂ Trend by Preset\n        fig_f2 = px.line(\n            df, x=df.index, y=\"f2\",\n            color=\"preset\", symbol=\"model\", markers=True,\n            title=\"F₂ Evolution Across Runs (Grouped by Preset)\",\n            labels={\"index\": \"Run #\", \"f2\": \"F₂ Score\"}\n        )\n        st.plotly_chart(fig_f2, use_container_width=True)\n\n        # 📊 Precision vs Recall Trend\n        fig_pr = px.line(\n            df, x=df.index, y=[\"precision\", \"recall\"],\n            color_discrete_sequence=[\"#007bff\", \"#ff5733\"],\n            title=\"Precision vs Recall Progression\",\n            labels={\"index\": \"Run #\", \"value\": \"Score\"},\n            markers=True\n        )\n        st.plotly_chart(fig_pr, use_container_width=True)\n\n        # 📊 F₂ Improvement Heatmap\n        fig_delta = px.bar(\n            df, x=df.index, y=\"ΔF2_%\", color=\"preset\",\n            title=\"ΔF₂ % Change Between Consecutive Runs\",\n            labels={\"index\": \"Run #\", \"ΔF2_%\": \"ΔF₂ (%)\"},\n            text=df[\"ΔF2_%\"].round(2)\n        )\n        fig_delta.update_traces(textposition=\"outside\")\n        st.plotly_chart(fig_delta, use_container_width=True)\n\n        # Summary Table\n        st.write(\"📋 Summary of Improvement Metrics:\")\n        delta_cols = [\"model\", \"preset\", \"f2\", \"precision\", \"recall\", \"ΔF2_%\", \"ΔRecall_%\", \"ΔPrecision_%\", \"train_time_s\", \"pred_count\"]\n        st.dataframe(df[delta_cols].round(3))\n\n        # Best Run Highlight\n        best_run = df.loc[df[\"f2\"].idxmax()]\n        st.success(\n            f\"🏆 Best F₂ = {best_run['f2']:.3f} \"\n            f\"| Recall = {best_run['recall']:.3f} \"\n            f\"| Precision = {best_run['precision']:.3f} \"\n            f\"| Preset: {best_run.get('preset','balanced')} \"\n            f\"| Model: {best_run.get('model','Tiny3DCNN')}\"\n        )\n\n    else:\n        st.info(\"Need at least 2 logged runs to visualize ΔF₂ trends.\")\nelse:\n    st.info(\"No benchmark_runs.csv found — run the model at least twice to generate F₂ deltas.\")\n#************************\n## Part 3 – Ablations\n#************************\nimport itertools, json, time, gc, torch\nfrom datetime import datetime\nfrom pathlib import Path\nimport pandas as pd\nimport numpy as np\nimport plotly.express as px\nimport streamlit as st\n\n# ---------------- Sidebar UI ----------------\nst.sidebar.header(\"⚙️ Ablation Controller\")\n\nwith st.sidebar.expander(\"Sweep Settings\", expanded=True):\n    patch_sizes = st.multiselect(\"Patch sizes (voxels)\", [32, 48, 64, 80], default=[64])\n    stride_factors = st.multiselect(\"Stride factors (speed tradeoff)\", [0.75, 1.0, 1.25, 1.5], default=[1.0])\n    dbscan_eps_list = st.multiselect(\"DBSCAN eps (voxels)\", [1.0, 1.5, 2.0, 3.0], default=[1.5])\n    presets = st.multiselect(\"Training preset\", [\"fast\", \"balanced\", \"thorough\"], default=[\"balanced\"])\n    max_runs = st.number_input(\"Max total runs (safety cap)\", min_value=1, max_value=200, value=10)\n    max_runtime_per_run = st.slider(\"Max runtime per run (minutes)\", 5, 120, 30)\n    max_total_runtime = st.number_input(\"Max total runtime (minutes)\", 10, 480, 120)\n\nmodel_choice = st.sidebar.selectbox(\"Select Architecture\", options=[\"tinycnn\"], index=0)\n\nrun_button = st.sidebar.button(\"🚀 Run Ablation Sweep\")\nsave_button = st.sidebar.button(\"💾 Save Current Results\")\n\n# ---------------- Session State ----------------\nif \"ablation_results\" not in st.session_state:\n    st.session_state[\"ablation_results\"] = []\n\n# ---------------- Helpers ----------------\ndef _make_grid(patches, strides, eps_list, presets_list, cap):\n    combos = list(itertools.product(patches, strides, eps_list, presets_list))\n    return combos[:cap]\n\ndef _save_results(results, prefix):\n    if not results:\n        return None\n    Path(CACHE_DIR).mkdir(parents=True, exist_ok=True)\n    ts = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n    csvp = Path(CACHE_DIR) / f\"{prefix}_{ts}.csv\"\n    jsonp = Path(CACHE_DIR) / f\"{prefix}_{ts}.json\"\n    pd.DataFrame(results).to_csv(csvp, index=False)\n    with open(jsonp, \"w\") as f:\n        json.dump(results, f, indent=2)\n    return str(csvp), str(jsonp)\n\ndef _run_single_combo(volume_memmap, gt_points, patch_size, stride_factor, dbscan_eps, preset, model_name, max_runtime_s):\n    t0 = time.time()\n    global DBSCAN_EPS\n    prev_eps = DBSCAN_EPS\n    DBSCAN_EPS = float(dbscan_eps)\n\n    try:\n        _, metrics = run_model_pipeline(\n            volume_memmap_or_arr=volume_memmap,\n            gt_points=gt_points,\n            model_name=model_name,\n            epochs=EPOCHS,\n            preset=preset,\n            patch_size=int(patch_size),\n            stride_factor=float(stride_factor)\n        )\n    finally:\n        DBSCAN_EPS = prev_eps\n\n    if time.time() - t0 > max_runtime_s:\n        raise TimeoutError(\"Run exceeded time budget\")\n\n    if gt_points.shape[0] == 0:\n        metrics = {\"precision\": np.nan, \"recall\": np.nan, \"f2\": np.nan, \"mean_dist_nm\": np.nan, \"note\": \"No GT available\"}\n\n    metrics.update({\n        \"model\": model_name,\n        \"patch_size\": int(patch_size),\n        \"stride_factor\": float(stride_factor),\n        \"dbscan_eps\": float(dbscan_eps),\n        \"preset\": preset,\n        \"run_time_s\": round(time.time() - t0, 2),\n        \"timestamp\": datetime.now().isoformat(),\n    })\n\n    gc.collect()\n    torch.cuda.empty_cache()\n    return metrics\n\n# ---------------- Results ----------------\nst.header(\"📊 Ablation Results Summary (nm-normalized)\")\nif st.session_state[\"ablation_results\"]:\n    df = pd.DataFrame(st.session_state[\"ablation_results\"])\n    if \"f2\" not in df.columns:\n        st.info(\"No successful runs yet — metrics not available.\")\n        st.dataframe(df)\n    else:\n        if \"mean_dist_nm\" in df.columns:\n            df[\"weighted_f2\"] = df[\"f2\"] / (1 + df[\"mean_dist_nm\"].clip(lower=1))\n        else:\n            df[\"weighted_f2\"] = df[\"f2\"]\n        df = df.sort_values(\"weighted_f2\", ascending=False)\n        st.dataframe(df)\n        fig = px.bar(df.head(10), y=\"weighted_f2\", hover_data=[\"model\", \"patch_size\", \"stride_factor\", \"dbscan_eps\", \"preset\", \"precision\", \"recall\"], title=\"Top Ablation Runs by Weighted F₂\")\n        st.plotly_chart(fig, use_container_width=True)\nelse:\n    st.info(\"No ablation results yet.\")\n\n# ---------------- Manual Save ----------------\nif save_button and st.session_state[\"ablation_results\"]:\n    paths = _save_results(st.session_state[\"ablation_results\"], \"ablation_manual\")\n    if paths:\n        st.success(f\"Saved to {paths[0]} and {paths[1]}\")\n\n\n# ===============================================================\n# Part 4 — Multi-Architecture Integration (T4-Optimized, Unified)\n# ===============================================================\nimport math\nimport torch.nn.functional as F\n\n# ---- Lightweight Squeeze-Excite (cheap attention) ----\nclass ConvSEBlock3D(nn.Module):\n    \"\"\"Small conv block with squeeze-excite channel attention (3D).\"\"\"\n    def __init__(self, in_ch, out_ch, stride=1):\n        super().__init__()\n        self.conv = nn.Conv3d(in_ch, out_ch, kernel_size=3, padding=1, stride=stride, bias=False)\n        self.bn = nn.BatchNorm3d(out_ch)\n        self.relu = nn.ReLU(inplace=True)\n        # SE\n        self.se_fc1 = nn.Conv3d(out_ch, max(1, out_ch // 8), kernel_size=1)\n        self.se_fc2 = nn.Conv3d(max(1, out_ch // 8), out_ch, kernel_size=1)\n\n    def forward(self, x):\n        out = self.relu(self.bn(self.conv(x)))\n        se = out.mean(dim=(2, 3, 4), keepdim=True)\n        se = F.relu(self.se_fc1(se))\n        se = torch.sigmoid(self.se_fc2(se))\n        return out * se\n\n\n# ---- Tiny 3D CNN (baseline, fastest) ----\nclass Tiny3DCNN(nn.Module):\n    def __init__(self, in_ch=1, base_ch=16):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv3d(in_ch, base_ch, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(base_ch, base_ch * 2, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(base_ch * 2, base_ch, 3, padding=1), nn.ReLU(),\n            nn.Conv3d(base_ch, 1, 1), nn.Sigmoid()\n        )\n    def forward(self, x):\n        return self.net(x)\n\n\n# ---- Small 3D U-Net (memory-conscious) ----\nclass UNet3D_small(nn.Module):\n    \"\"\"Depth = 3; base_ch controls width (8–32 recommended).\"\"\"\n    def __init__(self, in_ch=1, base_ch=16):\n        super().__init__()\n        self.enc1 = nn.Sequential(ConvSEBlock3D(in_ch, base_ch),\n                                  ConvSEBlock3D(base_ch, base_ch))\n        self.pool1 = nn.MaxPool3d(2)\n        self.enc2 = nn.Sequential(ConvSEBlock3D(base_ch, base_ch * 2),\n                                  ConvSEBlock3D(base_ch * 2, base_ch * 2))\n        self.pool2 = nn.MaxPool3d(2)\n        self.bottleneck = nn.Sequential(ConvSEBlock3D(base_ch * 2, base_ch * 4),\n                                        ConvSEBlock3D(base_ch * 4, base_ch * 4))\n        # decoder\n        self.up2 = nn.ConvTranspose3d(base_ch * 4, base_ch * 2, 2, stride=2)\n        self.dec2 = nn.Sequential(ConvSEBlock3D(base_ch * 4, base_ch * 2),\n                                  ConvSEBlock3D(base_ch * 2, base_ch * 2))\n        self.up1 = nn.ConvTranspose3d(base_ch * 2, base_ch, 2, stride=2)\n        self.dec1 = nn.Sequential(ConvSEBlock3D(base_ch * 2, base_ch),\n                                  ConvSEBlock3D(base_ch, base_ch))\n        self.final = nn.Conv3d(base_ch, 1, 1)\n        self.out_act = nn.Sigmoid()\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        b = self.bottleneck(self.pool2(e2))\n        d2 = self.dec2(torch.cat([self.up2(b), e2], 1))\n        d1 = self.dec1(torch.cat([self.up1(d2), e1], 1))\n        return self.out_act(self.final(d1))\n\n\n# ---- Small 3D ResNet ----\nclass ResBlock3D(nn.Module):\n    def __init__(self, in_ch, out_ch, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_ch, out_ch, 3, padding=1, stride=stride, bias=False)\n        self.bn1 = nn.BatchNorm3d(out_ch)\n        self.conv2 = nn.Conv3d(out_ch, out_ch, 3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm3d(out_ch)\n        self.relu = nn.ReLU(inplace=True)\n        self.down = None\n        if stride != 1 or in_ch != out_ch:\n            self.down = nn.Sequential(\n                nn.Conv3d(in_ch, out_ch, 1, stride=stride, bias=False),\n                nn.BatchNorm3d(out_ch)\n            )\n\n    def forward(self, x):\n        identity = x\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        if self.down is not None:\n            identity = self.down(identity)\n        out = self.relu(out + identity)\n        return out\n\n\nclass ResNet3D_small(nn.Module):\n    def __init__(self, in_ch=1, base_ch=16):\n        super().__init__()\n        self.inp = nn.Conv3d(in_ch, base_ch, 3, padding=1)\n        self.l1 = ResBlock3D(base_ch, base_ch)\n        self.pool1 = nn.MaxPool3d(2)\n        self.l2 = ResBlock3D(base_ch, base_ch * 2)\n        self.pool2 = nn.MaxPool3d(2)\n        self.l3 = ResBlock3D(base_ch * 2, base_ch * 4)\n        self.up2 = nn.ConvTranspose3d(base_ch * 4, base_ch * 2, 2, stride=2)\n        self.up1 = nn.ConvTranspose3d(base_ch * 2, base_ch, 2, stride=2)\n        self.outc = nn.Conv3d(base_ch, 1, 1)\n        self.act = nn.Sigmoid()\n\n    def forward(self, x):\n        x0 = F.relu(self.inp(x))\n        x1 = self.l1(x0)\n        x2 = self.l2(self.pool1(x1))\n        b = self.l3(self.pool2(x2))\n        out = self.up1(self.up2(b) + x2) + x1\n        return self.act(self.outc(out))\n\n\n# ---- Model factory ----\ndef get_model(name: str = \"tinycnn\", in_ch: int = 1, base_ch: int = 16, **kwargs):\n    \"\"\"Instantiate architecture by name.\"\"\"\n    n = name.lower()\n    if n in (\"tiny\", \"tinycnn\"):\n        return Tiny3DCNN(in_ch, base_ch)\n    if n in (\"unet\", \"unet3d\"):\n        return UNet3D_small(in_ch, base_ch)\n    if n in (\"resnet\", \"resnet3d\"):\n        return ResNet3D_small(in_ch, base_ch)\n    raise ValueError(f\"Unknown model name: {name}\")\n\n\n# ---- Unified runner (ready for ablation + Streamlit) ----\ndef run_model_pipeline(\n    volume_memmap_or_arr,\n    gt_points,\n    model: nn.Module = None,\n    model_name: str = \"tinycnn\",\n    model_kwargs: dict = None,\n    epochs: int = EPOCHS,\n    preset: str = \"balanced\",\n    patch_size: int = None,\n    stride_factor: float = 1.0,\n    base_ch: int = None,\n):\n    \"\"\"\n    Unified training/inference for all architectures.\n    Returns: clustered_points (Nx3 float), metrics (dict)\n    - Accepts either `model` (nn.Module) or `model_name` + `model_kwargs`.\n    - patch_size overrides global PATCH_SIZE if provided.\n    - stride_factor scales PATCH_STRIDE during inference.\n    - base_ch overrides model width if provided.\n    \"\"\"\n    import time, traceback, gc\n    t0 = time.time()\n    model_kwargs = dict(model_kwargs or {})\n    saved_path = None\n\n    # apply local patch/stride if provided (restore at end)\n    global PATCH_SIZE as _PATCH_SIZE_GLOBAL, PATCH_STRIDE as _PATCH_STRIDE_GLOBAL\n    prev_patch, prev_stride = PATCH_SIZE, PATCH_STRIDE\n    if patch_size is not None:\n        PATCH_SIZE = int(patch_size)\n        PATCH_STRIDE = max(1, int(PATCH_SIZE // 2))\n    else:\n        patch_size = PATCH_SIZE\n\n    try:\n        # Preset tuning\n        p = (preset or \"balanced\").lower()\n        if p == \"fast\":\n            epochs = max(1, int(epochs * 0.4))\n            patches_per_epoch = max(8, TRAIN_PATCHES_PER_EPOCH // 2)\n        elif p in (\"accurate\", \"thorough\"):\n            epochs = min(20, int(epochs * 1.5))\n            patches_per_epoch = max(TRAIN_PATCHES_PER_EPOCH, TRAIN_PATCHES_PER_EPOCH * 2)\n        else:\n            patches_per_epoch = TRAIN_PATCHES_PER_EPOCH\n\n        # Load volume\n        if isinstance(volume_memmap_or_arr, (str, Path)):\n            mm = np.load(volume_memmap_or_arr, mmap_mode=\"r\")\n            vol = np.array(mm)\n        else:\n            vol = np.array(volume_memmap_or_arr)\n        if vol.size == 0:\n            st.warning(\"Empty volume — cannot run model.\")\n            return np.zeros((0, 3)), {}\n\n        # Build target heatmap\n        Z, Y, X = vol.shape\n        target = np.zeros_like(vol, dtype=np.float32)\n        if gt_points is not None and len(gt_points) > 0:\n            for z, y, x in np.asarray(gt_points).astype(int):\n                if 0 <= z < Z and 0 <= y < Y and 0 <= x < X:\n                    target[z, y, x] = 1.0\n            from scipy import ndimage\n            target = ndimage.gaussian_filter(target, sigma=SIGMA_HEATMAP)\n        else:\n            target = np.zeros_like(vol, dtype=np.float32)\n\n        # Instantiate model (if not provided)\n        if model is None:\n            if base_ch is not None:\n                model_kwargs.setdefault(\"base_ch\", int(base_ch))\n            model = get_model(model_name, in_ch=1, **model_kwargs).to(DEVICE)\n        else:\n            model = model.to(DEVICE)\n\n        optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n        loss_fn = nn.MSELoss()\n        \n        # Determine scaler usage depending on the PyTorch version / environment\n        try:\n            scaler = torch.amp.GradScaler('cuda', enabled=(DEVICE == \"cuda\"))\n        except Exception:\n            scaler = torch.cuda.amp.GradScaler(enabled=(DEVICE == \"cuda\"))\n\n        total_steps = max(1, epochs * patches_per_epoch)\n        step = 0\n        pb = st.progress(0)\n        losses = []\n\n        st.info(f\"Training {model_name.upper()} ({p}) for {epochs} epoch(s) — patches/epoch={patches_per_epoch}\")\n\n        # Training loop (memory-friendly, patch-wise)\n        for ep in range(epochs):\n            patches_v, patches_t = extract_random_patches(vol, target, patches_per_epoch, PATCH_SIZE)\n            model.train()\n            for i in range(patches_v.shape[0]):\n                inp = torch.from_numpy(patches_v[i:i+1]).unsqueeze(1).to(DEVICE)\n                tgt = torch.from_numpy(patches_t[i:i+1]).unsqueeze(1).to(DEVICE)\n                optimizer.zero_grad()\n                \n                try:\n                    with torch.amp.autocast('cuda', enabled=(DEVICE == \"cuda\")):\n                        out = model(inp)\n                        loss = loss_fn(out, tgt)\n                except Exception:\n                    with torch.cuda.amp.autocast(enabled=(DEVICE == \"cuda\")):\n                        out = model(inp)\n                        loss = loss_fn(out, tgt)\n                        \n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n                losses.append(float(loss.detach().cpu().numpy()))\n                step += 1\n                if step % 5 == 0:\n                    pb.progress(int(100 * step / total_steps))\n                # free ASAP\n                del inp, tgt, out, loss\n                torch.cuda.empty_cache()\n            gc.collect(); torch.cuda.empty_cache()\n\n        # Inference (sliding window)\n        st.info(\"Running sliding-window inference...\")\n        recon, conf_map = sliding_window_inference(vol, model, patch_size=PATCH_SIZE, stride=int(PATCH_STRIDE * stride_factor))\n\n        # Threshold + clustering (FIXED: Coordinate Swap & Sigmoid Threshold)\n        # 1. Lowered hard minimum to 0.15 so we don't accidentally erase soft peaks\n        dynamic_thresh = (recon.mean() + 3.0 * recon.std()) * (1 - 0.15 * conf_map.mean())\n        thresh = max(0.15, dynamic_thresh) \n        \n        coords = np.array(np.nonzero(recon > thresh)).T\n        \n        if coords.shape[0] > MAX_POINTS:\n            idx = np.random.choice(coords.shape[0], size=MAX_POINTS, replace=False)\n            coords = coords[idx]\n            \n        coords_zyx = coords.astype(float)\n        \n        # 🧠 THE MASSIVE FIX: Extract CENTROIDS, not thousands of raw voxels!\n        if coords_zyx.shape[0] > 0:\n            from sklearn.cluster import DBSCAN\n            clustering = DBSCAN(eps=DBSCAN_EPS, min_samples=2).fit(coords_zyx)\n            labels = clustering.labels_\n            centroids = []\n            # Calculate the mean (center) of each distinct cluster\n            for k in set(labels):\n                if k != -1:  # Ignore unclustered noise (-1)\n                    centroids.append(coords_zyx[labels == k].mean(axis=0))\n            clustered = np.vstack(centroids) if centroids else np.zeros((0, 3))\n        else:\n            clustered = np.zeros((0, 3))\n\n        # 🧠 METRICS FIX: Safely use compute_metrics_nm to avoid the tuple crash\n        try:\n            if \"get_default_voxel_size\" in globals() and \"compute_metrics_nm\" in globals():\n                vx = get_default_voxel_size(\"byu_motor\")\n                # ✅ FIX: Relax the strictness radius to 40nm to allow for realistic motor volume sizes\n                metrics = compute_metrics_nm(clustered, gt_points, voxel_size_nm=vx, radius_nm=40.0)\n            else:\n                metrics = compute_metrics(clustered, gt_points, radius=5.0)\n        except Exception as e:\n            st.warning(f\"Metric computation warning: {e}\")\n            metrics = compute_metrics(clustered, gt_points, radius=5.0)\n            metrics[\"mean_dist_nm\"] = None\n\n        metrics.update({\n            \"model_name\": model_name if model_name else model.__class__.__name__,\n            \"patch_size\": int(PATCH_SIZE),\n            \"stride_factor\": float(stride_factor),\n            \"patches_per_epoch\": int(patches_per_epoch),\n            \"epochs\": int(epochs),\n            \"train_steps\": int(total_steps),\n            \"train_loss_mean\": float(np.mean(losses)) if losses else None,\n            \"train_time_s\": round(time.time() - t0, 2),\n            \"gpu_mem_MB\": round(torch.cuda.max_memory_allocated() / (1024**2), 2) if DEVICE == \"cuda\" else 0.0\n        })\n\n        # Save predictions (compressed)\n        try:\n            out_path = Path(CACHE_DIR) / f\"pred_{model_name}_{int(np.random.rand()*1e9)}.npz\"\n            np.savez_compressed(out_path, points=clustered)\n            saved_path = str(out_path)\n            metrics[\"saved_path\"] = saved_path\n        except Exception:\n            metrics[\"saved_path\"] = None\n\n        st.success(f\"✅ {model_name.upper()} done in {time.time() - t0:.1f}s — {len(clustered)} pts\")\n        gc.collect(); torch.cuda.empty_cache()\n        return clustered, metrics\n\n    except Exception as e:\n        # return empty results + error info\n        tb = traceback.format_exc()\n        st.error(f\"run_model_pipeline failed: {e}\")\n        return np.zeros((0, 3)), {\"error\": str(e), \"traceback\": tb}\n\n    finally:\n        # restore globals\n        PATCH_SIZE = prev_patch\n        PATCH_STRIDE = prev_stride\n        gc.collect(); torch.cuda.empty_cache()\n\n# ============================================================\n# Part 5 — Automated Reports & Visualization Dashboard (T4-safe, nm-aware, PDF Snapshot + Runtime Chart)\n# ============================================================\n\nimport io, time, psutil, platform\nfrom pathlib import Path\nfrom reportlab.lib.pagesizes import letter\nfrom reportlab.platypus import SimpleDocTemplate, Paragraph, Spacer, Table, TableStyle, Image\nfrom reportlab.lib import colors\nfrom reportlab.lib.styles import getSampleStyleSheet\nimport plotly.express as px\nimport plotly.io as pio\n\n# 🔒 Ensure CACHE_DIR is always a Path\nCACHE_DIR = Path(\"/kaggle/working/cache_phase1\")\n\nst.title(\"🔬 Cryo-ET 3D Localization Benchmark Dashboard\")\n\ntomos = find_tomograms(DATA_ROOT)\nif len(tomos) == 0:\n    st.error(\"No tomogram folders found under dataset path.\")\n    st.stop()\n\n# Sidebar configuration\nwith st.sidebar:\n    st.header(\"⚙️ Configuration\")\n    selected_name = st.selectbox(\"Select Tomogram\", [t.name for t in tomos])\n    selected_file = [t for t in tomos if t.name == selected_name][0]\n    model_name = st.selectbox(\"Select Architecture\", [\"tinycnn\", \"unet3d\", \"resnet3d\"])\n    preset = st.radio(\"Preset\", [\"fast\", \"balanced\", \"accurate\"], index=1)\n    base_ch = st.slider(\"Base Channels\", 8, 32, 16, 4)\n    n_epochs = st.slider(\"Epochs\", 1, 10, 3)\n    run_trigger = st.button(\"🚀 Run Model\")\n\n# ---- Load data ----\nraw_memmap = load_volume_from_jpegs_cached(str(selected_file))\nproc_cache_path = CACHE_DIR / f\"{selected_file.name}_proc.npy\"\n\nif proc_cache_path.exists():\n    # If we already saved it specifically for this tomogram, load it safely\n    vol_memmap = np.load(proc_cache_path, mmap_mode=\"r\")\nelse:\n    # 1. Process it (Streamlit cache might strip the memmap properties here)\n    processed_arr = preprocess_volume_cached_memmap(raw_memmap)\n    \n    # 2. Explicitly save the array to our known path to guarantee it is on disk\n    np.save(proc_cache_path, np.array(processed_arr, dtype=np.float32))\n    \n    # 3. Reload it directly with NumPy so it is a 100% valid memmap\n    vol_memmap = np.load(proc_cache_path, mmap_mode=\"r\")\n\nst.write(f\"Volume shape: {vol_memmap.shape}, memmap: {proc_cache_path}\")\n\n# ---- Load GT ----\ngt_points = np.zeros((0, 3))\ngt_csv = Path(DATA_ROOT) / GT_CSV\nif gt_csv.exists():\n    try:\n        df = pd.read_csv(gt_csv)\n        sel = df[df[\"tomo_id\"] == selected_name]\n        # ✅ FIX: Load in 0, 1, 2 order to perfectly match numpy's Z, Y, X shape\n        if {\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"}.issubset(sel.columns):\n            gt_points = sel[[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values.astype(float)\n    except Exception:\n        pass\n# ---- Tabs ----\ntab_run, tab_metrics, tab_viz, tab_report = st.tabs([\"🏃 Run\", \"📈 Metrics\", \"🧭 Visualization\", \"📄 Report Export\"])\n\n# ------------------------------------------------\n# TAB 1 — Run\n# ------------------------------------------------\nwith tab_run:\n    if run_trigger:\n        st.info(f\"Running {model_name.upper()} ({preset})...\")\n        start_time = time.time()\n        mem_before = psutil.virtual_memory().used / (1024**3)\n\n        recon_points, metrics = run_model_pipeline(\n            vol_memmap, gt_points,\n            model_name=model_name,\n            model_kwargs={\"base_ch\": base_ch},\n            epochs=n_epochs,\n            preset=preset\n        )\n\n        runtime = time.time() - start_time\n        gpu_mem = torch.cuda.memory_allocated() / (1024**3) if DEVICE == \"cuda\" else 0.0\n\n        matched, gt_count, dists = match_predictions_to_gt(recon_points, gt_points)\n        results_row = {\n            \"tomo\": selected_name,\n            \"model\": model_name,\n            \"preset\": preset,\n            \"base_ch\": base_ch,\n            \"runtime_s\": round(runtime, 2),\n            \"gpu_mem_gb\": round(gpu_mem, 2),\n            \"matched\": matched,\n            \"gt_count\": int(gt_count),\n            \"pred_count\": len(recon_points),\n            \"mean_dist\": float(np.mean(dists)) if len(dists) else 0.0,\n            \"mean_dist_nm\": metrics.get(\"mean_dist_nm\", np.nan),\n            \"f2\": metrics.get(\"f2\", np.nan),\n            \"precision\": metrics.get(\"precision\", np.nan),\n            \"recall\": metrics.get(\"recall\", np.nan),\n            \"timestamp\": time.strftime(\"%Y-%m-%d %H:%M:%S\"),\n        }\n\n        metrics_path = CACHE_DIR / \"benchmark_metrics.csv\"\n        pd.DataFrame([results_row]).to_csv(metrics_path, mode=\"a\", index=False, header=not metrics_path.exists())\n\n        st.success(f\"✅ {model_name.upper()} done in {runtime:.1f}s, GPU {gpu_mem:.2f} GB used\")\n        st.write(f\"Matched {matched}/{gt_count}, mean distance = {results_row['mean_dist']:.3f}\")\n        st.write(f\"F₂ = {metrics.get('f2',0):.3f} Precision = {metrics.get('precision',0):.3f} Recall = {metrics.get('recall',0):.3f}\")\n        np.save(CACHE_DIR / f\"{selected_name}_{model_name}_pts.npy\", recon_points)\n    else:\n        st.info(\"Click 🚀 Run Model in sidebar to start training & inference.\")\n\n# ------------------------------------------------\n# TAB 2 — Metrics\n# ------------------------------------------------\nwith tab_metrics:\n    metrics_path = CACHE_DIR / \"benchmark_metrics.csv\"\n    if metrics_path.exists():\n        dfm = pd.read_csv(metrics_path)\n        st.dataframe(dfm.sort_values(\"timestamp\", ascending=False), use_container_width=True)\n\n        if len(dfm) > 1:\n            # 🧠 Ensure 'mean_dist_nm' column exists\n            if \"mean_dist_nm\" not in dfm.columns:\n                if \"mean_dist\" in dfm.columns:\n                    dfm[\"mean_dist_nm\"] = dfm[\"mean_dist\"]\n                else:\n                    dfm[\"mean_dist_nm\"] = 0.0\n\n            # ✅ Dynamically choose hover columns that exist\n            hover_cols = [c for c in [\"tomo\", \"pred_count\", \"f2\", \"precision\", \"recall\"]\n                          if c in dfm.columns]\n\n            fig = px.scatter_3d(\n                dfm,\n                x=\"runtime_s\",\n                y=\"mean_dist_nm\",\n                z=\"gpu_mem_gb\",\n                color=\"model\",\n                symbol=\"preset\",\n                size=\"matched\",\n                hover_data=hover_cols,\n                title=\"Runtime vs Accuracy (nm) vs GPU Usage\",\n            )\n            st.plotly_chart(fig, use_container_width=True)\n    else:\n        st.info(\"No benchmark metrics yet — run at least one model.\")\n\n\n# ------------------------------------------------\n# TAB 3 — Visualization\n# ------------------------------------------------\nwith tab_viz:\n    st.subheader(\"3D Mesh + Points (adaptive)\")\n    mesh_verts, mesh_faces = marching_cubes_mesh_adaptive(proc_cache_path)\n\n    pred_path = CACHE_DIR / f\"{selected_name}_{model_name}_pts.npy\"\n    recon_points = np.load(pred_path) if pred_path.exists() else np.zeros((0, 3))\n\n    fig3d = plot_volume_mesh_and_points(\n        mesh_verts, mesh_faces,\n        points=recon_points, gt_pts=gt_points, pred_pts=recon_points,\n        title=selected_name\n    )\n    st.plotly_chart(fig3d, use_container_width=True)\n\n    # 📸 Save snapshot for report\n    snapshot_path = CACHE_DIR / \"last_viz_snapshot.png\"\n    try:\n        pio.write_image(fig3d, str(snapshot_path), format=\"png\", width=900, height=700, scale=2)\n        st.success(f\"📸 Snapshot saved for report: {snapshot_path.name}\")\n    except Exception as e:\n        st.warning(f\"Snapshot export skipped: {e}\")\n\n    st.subheader(\"Slice Viewer\")\n    z_index = st.slider(\"Select Z slice\", 0, max(0, vol_memmap.shape[0]-1), 0)\n    fig_slice = plot_slice_with_points_memmap(\n        proc_cache_path, z_index,\n        points=recon_points, gt_pts=gt_points, pred_pts=recon_points,\n        title=selected_name\n    )\n    st.plotly_chart(fig_slice, use_container_width=True)\n\n# ------------------------------------------------\n# TAB 4 — Report Export (PDF with Visualization & F₂ Chart, Safe)\n# ------------------------------------------------\nwith tab_report:\n    st.subheader(\"📄 Generate Summary PDF (with Visualization & F₂ Chart)\")\n    metrics_path = CACHE_DIR / \"benchmark_metrics.csv\"\n    if metrics_path.exists():\n        dfm = pd.read_csv(metrics_path)\n\n        # 🧮 Auto-compute F₂ if missing and precision/recall available\n        if \"f2\" not in dfm.columns and {\"precision\", \"recall\"}.issubset(dfm.columns):\n            dfm[\"f2\"] = np.where(\n                (dfm[\"precision\"] + dfm[\"recall\"]) > 0,\n                5 * (dfm[\"precision\"] * dfm[\"recall\"]) /\n                (4 * dfm[\"precision\"] + dfm[\"recall\"]),\n                0\n            )\n\n        # 🔬 Add placeholder mean_dist_nm if absent\n        if \"mean_dist_nm\" not in dfm.columns:\n            dfm[\"mean_dist_nm\"] = dfm.get(\"mean_dist\", np.zeros(len(dfm)))\n\n        # 📊 Generate F₂ vs mean distance chart\n        chart_path = CACHE_DIR / \"runtime_vs_f2_chart.png\"\n        try:\n            if \"f2\" in dfm.columns and dfm[\"f2\"].notna().any():\n                fig_chart = px.scatter(\n                    dfm, x=\"mean_dist_nm\", y=\"f2\",\n                    color=\"model\", symbol=\"preset\", size=\"matched\",\n                    title=\"F₂ vs Mean Distance (nm)\",\n                    labels={\"mean_dist_nm\": \"Mean Distance (nm)\", \"f2\": \"F₂ Score\"}\n                )\n                import plotly.io as pio\n                try:\n                    pio.write_image(fig_chart, str(chart_path),\n                                    format=\"png\", width=800, height=600, scale=2)\n                    st.image(str(chart_path),\n                             caption=\"F₂ vs Mean Distance Chart (preview)\",\n                             use_container_width=True)\n                except Exception as e:\n                    st.warning(f\"Chart image export skipped (kaleido missing?): {e}\")\n            else:\n                st.info(\"No valid F₂ data available for chart rendering.\")\n        except Exception as e:\n            st.warning(f\"Chart export skipped: {e}\")\n\n        # 🧾 Assemble PDF\n        from reportlab.platypus import Image  # lazy import for safety\n        buf = io.BytesIO()\n        doc = SimpleDocTemplate(buf, pagesize=letter)\n        styles = getSampleStyleSheet()\n\n        best_f2 = dfm[\"f2\"].max() if \"f2\" in dfm.columns else 0.0\n        best_md = dfm[\"mean_dist_nm\"].min() if \"mean_dist_nm\" in dfm.columns else 0.0\n\n        story = [\n            Paragraph(\"<b>3D Cryo-ET Localization Benchmark Summary</b>\", styles[\"Title\"]),\n            Spacer(1, 12),\n            Paragraph(f\"System: {platform.node()} | {platform.processor()}\", styles[\"Normal\"]),\n            Paragraph(f\"Date: {time.strftime('%Y-%m-%d %H:%M:%S')}\", styles[\"Normal\"]),\n            Spacer(1, 12),\n            Paragraph(f\"Total Runs: {len(dfm)}\", styles[\"Normal\"]),\n            Paragraph(f\"Best F₂: {best_f2:.3f}\", styles[\"Normal\"]),\n            Paragraph(f\"Lowest Mean Distance (nm): {best_md:.2f}\", styles[\"Normal\"]),\n            Spacer(1, 12)\n        ]\n\n        # 🖼 Visualization snapshot & chart\n        snap_img = CACHE_DIR / \"last_viz_snapshot.png\"\n        if snap_img.exists():\n            story += [\n                Paragraph(\"<b>3D Visualization Snapshot</b>\", styles[\"Heading2\"]),\n                Spacer(1, 6),\n                Image(str(snap_img), width=450, height=350),\n                Spacer(1, 12)\n            ]\n        if chart_path.exists():\n            story += [\n                Paragraph(\"<b>F₂ vs Mean Distance Chart</b>\", styles[\"Heading2\"]),\n                Spacer(1, 6),\n                Image(str(chart_path), width=400, height=300),\n                Spacer(1, 12)\n            ]\n\n        # 📋 Metrics table\n        table_data = [list(dfm.columns)] + dfm.astype(str).values.tolist()\n        from reportlab.platypus import Table, TableStyle\n        tbl = Table(table_data, repeatRows=1)\n        tbl.setStyle(TableStyle([\n            (\"BACKGROUND\", (0, 0), (-1, 0), colors.grey),\n            (\"TEXTCOLOR\", (0, 0), (-1, 0), colors.whitesmoke),\n            (\"ALIGN\", (0, 0), (-1, -1), \"CENTER\"),\n            (\"GRID\", (0, 0), (-1, -1), 0.25, colors.black),\n        ]))\n        story.append(tbl)\n\n        doc.build(story)\n        pdf_bytes = buf.getvalue()\n        st.download_button(\"⬇️ Download PDF Report\",\n                           pdf_bytes,\n                           file_name=\"CryoET_Benchmark_Report.pdf\")\n    else:\n        st.info(\"No metrics logged yet — run a model to generate report.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:04.511380Z","iopub.execute_input":"2026-02-22T14:41:04.511650Z","iopub.status.idle":"2026-02-22T14:41:04.544502Z","shell.execute_reply.started":"2026-02-22T14:41:04.511621Z","shell.execute_reply":"2026-02-22T14:41:04.543640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!sed -i 's/global PATCH_SIZE as _PATCH_SIZE_GLOBAL, PATCH_STRIDE as _PATCH_STRIDE_GLOBAL/global PATCH_SIZE, PATCH_STRIDE/' /kaggle/working/app.py\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:04.546227Z","iopub.execute_input":"2026-02-22T14:41:04.546582Z","iopub.status.idle":"2026-02-22T14:41:04.674502Z","shell.execute_reply.started":"2026-02-22T14:41:04.546564Z","shell.execute_reply":"2026-02-22T14:41:04.673422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!grep -n \"global PATCH_SIZE\" /kaggle/working/app.py\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:04.675687Z","iopub.execute_input":"2026-02-22T14:41:04.676541Z","iopub.status.idle":"2026-02-22T14:41:04.798379Z","shell.execute_reply.started":"2026-02-22T14:41:04.676514Z","shell.execute_reply":"2026-02-22T14:41:04.797678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Part 6 — Data Augmentation & Semi-Synthetic Expansion\n# (CPU-safe, nm-aware, with logging & visualization)\n# ============================================================\n\nimport os, random, gc\nimport numpy as np\nfrom pathlib import Path\nfrom scipy.ndimage import gaussian_filter, zoom\nimport streamlit as st\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n# --- ensure cache & aug directories exist ---\nCACHE_DIR = Path(\"/kaggle/working/cache_phase1\")\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\nAUG_DIR = CACHE_DIR / \"augmented\"\nAUG_DIR.mkdir(parents=True, exist_ok=True)\nAUG_LOG = CACHE_DIR / \"augmentation_log.csv\"\n\n# ------------------------------------------------------------\n# 🔬 Semi-Synthetic Augmentation Function\n# ------------------------------------------------------------\ndef generate_augmented_variants(volume_memmap_or_arr, name_prefix: str, n_variants: int = 3):\n    \"\"\"\n    Generate semi-synthetic noisy / blurred / rescaled variants\n    to simulate realistic Cryo-ET tomogram variability.\n    Returns list of .npy paths + summary metadata.\n    \"\"\"\n    if isinstance(volume_memmap_or_arr, (str, Path)):\n        mm = np.load(volume_memmap_or_arr, mmap_mode='r')\n        base = np.array(mm, dtype=np.float32)\n    else:\n        base = np.array(volume_memmap_or_arr, dtype=np.float32)\n\n    variants, meta = [], []\n    base_mean, base_std = float(base.mean()), float(base.std())\n\n    for i in range(n_variants):\n        vol = base.copy()\n\n        # --- 1️⃣ Add Gaussian + optional Poisson noise ---\n        sigma_n = np.random.uniform(0.01, 0.1) * vol.std()\n        gauss = np.random.normal(0, sigma_n, vol.shape).astype(np.float32)\n        vol += gauss\n        if np.random.rand() < 0.5:\n            lam = np.random.uniform(0.5, 2.0)\n            vol = np.random.poisson(np.clip(vol, 0, None) * lam).astype(np.float32) / max(lam, 1e-6)\n\n        # --- 2️⃣ Defocus blur ---\n        blur_sigma = np.random.uniform(0.4, 1.5)\n        vol = gaussian_filter(vol, sigma=blur_sigma)\n\n        # --- 3️⃣ Random voxel scaling / downsampling ---\n        if np.random.rand() < 0.6:\n            scale = np.random.uniform(0.6, 1.0)\n            vol = zoom(vol, scale, order=1)\n            pad = [(max(0, b - v) // 2, max(0, b - v) - max(0, b - v) // 2)\n                   for b, v in zip(base.shape, vol.shape)]\n            vol = np.pad(vol, pad, mode='reflect')\n            vol = vol[:base.shape[0], :base.shape[1], :base.shape[2]]\n\n        # --- 4️⃣ Contrast normalization ---\n        vol = (vol - vol.mean()) / (vol.std() + 1e-8)\n        vol = np.clip(vol, -3, 3)\n\n        # --- 🧮 Compute SNR & save variant ---\n        snr = float(20 * np.log10((base_std + 1e-8) / (sigma_n + 1e-8)))\n        out_path = AUG_DIR / f\"{name_prefix}_aug{i+1}.npy\"\n        np.save(out_path, vol.astype(np.float32))\n        variants.append(out_path)\n\n        meta.append({\n            \"variant\": f\"{name_prefix}_aug{i+1}\",\n            \"gaussian_sigma\": round(sigma_n, 4),\n            \"blur_sigma\": round(blur_sigma, 3),\n            \"scale_factor\": round(scale if 'scale' in locals() else 1.0, 3),\n            \"snr_db\": round(snr, 2),\n            \"mean\": round(float(vol.mean()), 4),\n            \"std\": round(float(vol.std()), 4)\n        })\n\n        gc.collect()\n\n    # --- Save metadata log ---\n    df_meta = pd.DataFrame(meta)\n    if AUG_LOG.exists():\n        prev = pd.read_csv(AUG_LOG)\n        df_meta = pd.concat([prev, df_meta], ignore_index=True)\n    df_meta.to_csv(AUG_LOG, index=False)\n\n    return variants, df_meta\n\n\n# ------------------------------------------------------------\n# 🧭 Streamlit Integration\n# ------------------------------------------------------------\nwith st.expander(\"🧬 Semi-Synthetic Data Augmentation\", expanded=False):\n    st.markdown(\"\"\"\n    Create semi-synthetic tomograms to mimic realistic Cryo-ET variability.\n    Each variant introduces random blur, noise, voxel scaling, and normalization\n    to emulate electron dose, focus drift, and reconstruction artifacts.\n    \"\"\")\n\n    n_aug = st.slider(\"Number of augmented variants\", 1, 6, 3)\n    if st.button(\"Generate Augmented Tomograms\"):\n        with st.spinner(\"Generating semi-synthetic tomograms...\"):\n            aug_paths, meta_df = generate_augmented_variants(proc_cache_path, selected_name, n_variants=n_aug)\n        st.success(f\"✅ {len(aug_paths)} augmented tomograms saved to {AUG_DIR}\")\n        st.dataframe(meta_df)\n\n        # 🖼 Preview middle slice & histogram\n        rand_path = random.choice(aug_paths)\n        rand_aug = np.load(rand_path, mmap_mode='r')\n        mid_z = rand_aug.shape[0] // 2\n\n        col1, col2 = st.columns(2)\n        with col1:\n            st.image(rand_aug[mid_z, :, :], caption=f\"Preview of {rand_path.name}\", use_container_width=True)\n        with col2:\n            fig, ax = plt.subplots()\n            ax.hist(rand_aug.flatten(), bins=50, color='gray')\n            ax.set_title(\"Voxel Intensity Distribution\")\n            st.pyplot(fig)\n\n        st.caption(\"Tip: You can now re-run your TinyCNN or U-Net using these augmented volumes for data diversity.\")\n\n# ------------------------------------------------------------\n# Optional Hook: Direct Augment → Train Loop (for automation)\n# ------------------------------------------------------------\nif st.button(\"🧠 Augment & Train Automatically (TinyCNN baseline)\"):\n    st.info(\"Running augmentation + model pipeline...\")\n    aug_paths, _ = generate_augmented_variants(proc_cache_path, selected_name, n_variants=n_aug)\n    for path in aug_paths:\n        st.write(f\"Training on {path.name} ...\")\n        recon, metrics = run_model_pipeline(\n            path, gt_points,\n            model_name=\"tinycnn\",\n            epochs=3,\n            preset=\"fast\"\n        )\n        st.success(f\"Done: F₂={metrics.get('f2',0):.3f} | MeanDist_nm={metrics.get('mean_dist_nm',0):.2f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:04.799773Z","iopub.execute_input":"2026-02-22T14:41:04.800094Z","iopub.status.idle":"2026-02-22T14:41:06.317089Z","shell.execute_reply.started":"2026-02-22T14:41:04.800059Z","shell.execute_reply":"2026-02-22T14:41:06.316080Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# Part 7 — Reproducibility Tracker, Seed Control & Auto-Logging\n# ================================================================\n\nimport os, json, random, socket, subprocess, warnings, hashlib\nfrom datetime import datetime\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport platform\nimport streamlit as st\nimport plotly.express as px\n\n# ---------------- Directories ----------------\nCACHE_DIR = Path(\"/kaggle/working/cache_phase1\")\nLOG_DIR = CACHE_DIR / \"experiment_logs\"\nLOG_DIR.mkdir(exist_ok=True, parents=True)\nSUMMARY_CSV = CACHE_DIR / \"experiment_summary.csv\"\nAUG_LOG = CACHE_DIR / \"augmentation_log.csv\"\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# ---------------- Seed control ----------------\ndef set_seed(seed: int = 42, deterministic: bool = False):\n    \"\"\"Set global seeds for reproducibility and return RNG fingerprint.\"\"\"\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n    torch.backends.cudnn.deterministic = deterministic\n    torch.backends.cudnn.benchmark = not deterministic\n\n    # create reproducibility fingerprint\n    state = (\n        str(seed)\n        + str(np.random.get_state())\n        + str(random.getstate()[1][:10])\n        + str(torch.get_rng_state()[:10].tolist())\n    )\n    fingerprint = hashlib.sha1(state.encode()).hexdigest()[:12]\n    return seed, fingerprint\n\n\n# ✅ default seed + fingerprint\nGLOBAL_SEED, GLOBAL_FP = set_seed(42)\n\n# ---------------- Environment capture ----------------\ndef _get_env_info():\n    \"\"\"Collect environment + package version info.\"\"\"\n    info = {\n        \"python\": platform.python_version(),\n        \"hostname\": socket.gethostname(),\n        \"os\": platform.platform(),\n        \"seed\": GLOBAL_SEED,\n        \"fingerprint\": GLOBAL_FP,\n    }\n\n    pkgs = {}\n    for pkg in (\"torch\", \"numpy\", \"pandas\", \"scipy\", \"sklearn\", \"plotly\"):\n        try:\n            mod = __import__(pkg)\n            pkgs[pkg] = getattr(mod, \"__version__\", \"unknown\")\n        except Exception:\n            pkgs[pkg] = \"not_installed\"\n    info[\"packages\"] = pkgs\n\n    if torch.cuda.is_available():\n        props = torch.cuda.get_device_properties(0)\n        info[\"gpu\"] = {\n            \"name\": props.name,\n            \"total_mem_GB\": round(props.total_memory / 1e9, 2),\n            \"count\": torch.cuda.device_count(),\n        }\n    else:\n        info[\"gpu\"] = None\n\n    try:\n        commit = subprocess.check_output([\"git\", \"rev-parse\", \"HEAD\"], stderr=subprocess.DEVNULL).decode().strip()\n        info[\"git_commit\"] = commit\n    except Exception:\n        info[\"git_commit\"] = None\n    return info\n\n\n# ---------------- Checkpoint helpers ----------------\ndef save_checkpoint(model: nn.Module, optimizer=None, path: Path = None, extra: dict = None):\n    \"\"\"Save model weights (+ optimizer if provided).\"\"\"\n    if path is None:\n        path = LOG_DIR / f\"ckpt_{datetime.now().strftime('%Y%m%d_%H%M%S')}.pth\"\n    try:\n        payload = {\"model_state\": model.state_dict()}\n        if optimizer is not None:\n            payload[\"optimizer_state\"] = optimizer.state_dict()\n        if extra:\n            payload[\"extra\"] = extra\n        torch.save(payload, str(path))\n        return str(path)\n    except Exception as e:\n        warnings.warn(f\"Checkpoint save failed: {e}\")\n        return None\n\n\ndef load_checkpoint(path: str, model: nn.Module = None, optimizer=None, map_location=None):\n    \"\"\"Load checkpoint into model/optimizer.\"\"\"\n    chk = torch.load(path, map_location=(map_location or DEVICE))\n    if model is not None and \"model_state\" in chk:\n        model.load_state_dict(chk[\"model_state\"])\n    if optimizer is not None and \"optimizer_state\" in chk:\n        optimizer.load_state_dict(chk[\"optimizer_state\"])\n    return chk\n\n\n# ---------------- Experiment logging ----------------\ndef _make_run_id():\n    return datetime.now().strftime(\"%Y%m%d_%H%M%S\") + f\"_{np.random.randint(0,1e6):06d}\"\n\n\ndef log_experiment(config: dict, metrics: dict, model: nn.Module = None, optimizer=None, save_model: bool = False):\n    \"\"\"Atomically log an experiment run + optional checkpoint.\"\"\"\n    run_id = _make_run_id()\n    env = _get_env_info()\n    log = {\n        \"run_id\": run_id,\n        \"timestamp\": datetime.now().isoformat(),\n        \"config\": config,\n        \"metrics\": metrics,\n        \"env\": env,\n    }\n\n    json_path = LOG_DIR / f\"exp_{run_id}.json\"\n    with open(json_path, \"w\") as f:\n        json.dump(log, f, indent=2)\n\n    model_path = None\n    if save_model and model is not None:\n        model_path = save_checkpoint(\n            model, optimizer=optimizer,\n            path=LOG_DIR / f\"model_{run_id}.pth\",\n            extra={\"run_id\": run_id, \"fingerprint\": GLOBAL_FP},\n        )\n\n    # flat summary CSV\n    row = {\n        \"run_id\": run_id,\n        \"timestamp\": log[\"timestamp\"],\n        \"tomo\": config.get(\"tomo\"),\n        \"model\": config.get(\"model_name\", config.get(\"model\")),\n        \"preset\": config.get(\"preset\"),\n        \"mean_dist\": metrics.get(\"mean_dist\"),\n        \"precision\": metrics.get(\"precision\"),\n        \"recall\": metrics.get(\"recall\"),\n        \"f2\": metrics.get(\"f2\"),\n        \"runtime_s\": metrics.get(\"runtime_s\"),\n        \"gpu_mem_gb\": metrics.get(\"gpu_mem_gb\"),\n        \"seed\": env.get(\"seed\"),\n        \"fingerprint\": env.get(\"fingerprint\"),\n        \"model_path\": model_path,\n    }\n\n    df_row = pd.DataFrame([row])\n    df_row.to_csv(SUMMARY_CSV, mode=\"a\", index=False, header=not SUMMARY_CSV.exists())\n\n    # link augmentation metadata if present\n    if AUG_LOG.exists():\n        df_aug = pd.read_csv(AUG_LOG)\n        df_aug[\"linked_run_id\"] = run_id\n        df_aug.to_csv(AUG_LOG, index=False)\n\n    return {\"json\": str(json_path), \"model\": model_path, \"summary\": row}\n\n\n# ---------------- Load all logs ----------------\ndef load_all_experiments(as_df: bool = True):\n    rows = []\n    for p in sorted(LOG_DIR.glob(\"exp_*.json\")):\n        try:\n            with open(p, \"r\") as f:\n                rows.append(json.load(f))\n        except Exception:\n            continue\n    if as_df:\n        flat = []\n        for r in rows:\n            flat.append({\n                \"run_id\": r.get(\"run_id\"),\n                \"timestamp\": r.get(\"timestamp\"),\n                \"tomo\": r.get(\"config\", {}).get(\"tomo\"),\n                \"model\": r.get(\"config\", {}).get(\"model_name\", r.get(\"config\", {}).get(\"model\")),\n                \"preset\": r.get(\"config\", {}).get(\"preset\"),\n                \"mean_dist\": r.get(\"metrics\", {}).get(\"mean_dist\"),\n                \"precision\": r.get(\"metrics\", {}).get(\"precision\"),\n                \"recall\": r.get(\"metrics\", {}).get(\"recall\"),\n                \"f2\": r.get(\"metrics\", {}).get(\"f2\"),\n                \"seed\": r.get(\"env\", {}).get(\"seed\"),\n                \"fingerprint\": r.get(\"env\", {}).get(\"fingerprint\"),\n            })\n        return pd.DataFrame(flat)\n    else:\n        return rows\n\n\n# ---------------- Streamlit UI ----------------\nwith st.expander(\"🗂️ Experiment Logs & Reproducibility\", expanded=False):\n    st.write(\"All experiment JSON logs are saved in:\", str(LOG_DIR))\n    if SUMMARY_CSV.exists():\n        df_summary = pd.read_csv(SUMMARY_CSV)\n        st.dataframe(df_summary.sort_values(\"timestamp\", ascending=False))\n\n        # 📊 Quick visual summary\n        if len(df_summary) > 1:\n            fig = px.scatter(\n                df_summary,\n                x=\"runtime_s\", y=\"f2\", color=\"model\",\n                size=\"precision\", hover_data=[\"tomo\", \"recall\", \"mean_dist\"],\n                title=\"Experiment F₂ vs Runtime (per model)\"\n            )\n            st.plotly_chart(fig, use_container_width=True)\n\n        st.download_button(\n            \"⬇️ Download CSV\", df_summary.to_csv(index=False).encode(), \"experiment_summary.csv\"\n        )\n    else:\n        st.info(\"No summary CSV yet — run experiments and call log_experiment(...) to create entries.\")\n\n    if st.button(\"Reload JSON Logs\"):\n        df_logs = load_all_experiments(as_df=True)\n        if df_logs.empty:\n            st.info(\"No JSON logs found.\")\n        else:\n            st.dataframe(df_logs.sort_values(\"timestamp\", ascending=False))\n# ------------------------------------------------------------\n# 🔬 Experiment Comparison Utility\n# ------------------------------------------------------------\nst.markdown(\"---\")\nst.subheader(\"🔍 Compare Two Experiments\")\n\nif SUMMARY_CSV.exists():\n    df_summary = pd.read_csv(SUMMARY_CSV)\n    run_ids = df_summary[\"run_id\"].tolist()\n    if len(run_ids) >= 2:\n        col1, col2 = st.columns(2)\n        with col1:\n            run_a = st.selectbox(\"Select First Run\", run_ids, key=\"cmp_a\")\n        with col2:\n            run_b = st.selectbox(\"Select Second Run\", run_ids, index=1, key=\"cmp_b\")\n\n        if st.button(\"Compare Selected Runs\"):\n            row_a = df_summary[df_summary[\"run_id\"] == run_a].iloc[0].to_dict()\n            row_b = df_summary[df_summary[\"run_id\"] == run_b].iloc[0].to_dict()\n\n            # --- Compute differences ---\n            metrics = [\"mean_dist\", \"precision\", \"recall\", \"f2\", \"runtime_s\"]\n            diff = {m: round(row_b[m] - row_a[m], 4) for m in metrics if m in row_a}\n\n            # --- Display tables ---\n            st.markdown(\"#### 🧾 Configuration Comparison\")\n            cfg_df = pd.DataFrame([\n                {\"Parameter\": k, \"Run A\": row_a.get(k), \"Run B\": row_b.get(k)}\n                for k in [\"tomo\", \"model\", \"preset\", \"seed\", \"fingerprint\", \"gpu_mem_gb\"]\n            ])\n            st.dataframe(cfg_df, use_container_width=True)\n\n            st.markdown(\"#### 📊 Metric Deltas (B − A)\")\n            delta_df = pd.DataFrame(\n                [{\"Metric\": k, \"Δ Value\": v, \"Direction\": \"⬆️\" if v > 0 else \"⬇️\"} for k, v in diff.items()]\n            )\n            st.dataframe(delta_df, use_container_width=True)\n\n            # --- Optional similarity score ---\n            match_keys = [\"tomo\", \"model\", \"preset\", \"seed\"]\n            match_score = sum(row_a[k] == row_b[k] for k in match_keys) / len(match_keys)\n            st.info(f\"🔗 Configuration similarity score: **{match_score:.2f}**\")\n\n            # --- Visualization ---\n            import plotly.graph_objects as go\n            fig = go.Figure()\n            for k in metrics:\n                if k in row_a and k in row_b:\n                    fig.add_trace(go.Bar(\n                        x=[k],\n                        y=[row_a[k]],\n                        name=f\"{run_a} (A)\",\n                        marker_color=\"royalblue\"\n                    ))\n                    fig.add_trace(go.Bar(\n                        x=[k],\n                        y=[row_b[k]],\n                        name=f\"{run_b} (B)\",\n                        marker_color=\"darkorange\"\n                    ))\n            fig.update_layout(barmode=\"group\", title=\"Metric Comparison (A vs B)\", height=400)\n            st.plotly_chart(fig, use_container_width=True)\n    else:\n        st.info(\"Need at least two experiment logs to compare.\")\nelse:\n    st.warning(\"No experiments logged yet — run a model to enable comparison.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:06.318279Z","iopub.execute_input":"2026-02-22T14:41:06.318847Z","iopub.status.idle":"2026-02-22T14:41:11.540219Z","shell.execute_reply.started":"2026-02-22T14:41:06.318826Z","shell.execute_reply":"2026-02-22T14:41:11.539492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile -a /kaggle/working/app.py\n\n# ================================================================\n# 🚀 MAIN RUNNER FOR ABLATION (Must be at the very bottom of file)\n# ================================================================\nif run_button:\n    tomos = find_tomograms(DATA_ROOT)\n    if not tomos:\n        st.error(\"❌ No tomograms found.\")\n        st.stop()\n\n    tomo_names = [t.name for t in tomos]\n    selected_tomo_name = st.selectbox(\"Select Tomogram for Ablation\", tomo_names)\n    selected_tomo = next((t for t in tomos if t.name == selected_tomo_name), None)\n    \n    if selected_tomo is None:\n        st.error(\"Tomogram not found.\")\n        st.stop()\n\n    try:\n        raw_memmap = load_volume_from_jpegs_cached(selected_tomo)\n        if raw_memmap is None:\n            raise RuntimeError(\"Raw volume load failed\")\n        volume_memmap = preprocess_volume_cached_memmap(raw_memmap)\n        volume_memmap = ensure_disk_memmap(volume_memmap)\n        if not isinstance(volume_memmap, np.memmap) or volume_memmap.filename is None:\n            raise RuntimeError(\"volume_memmap is not disk-backed\")\n    except Exception as e:\n        st.error(f\"❌ Failed to load tomogram: {e}\")\n        st.stop()\n\n    gt_points = np.zeros((0, 3))\n    gt_csv = Path(DATA_ROOT) / GT_CSV\n    if gt_csv.exists():\n        df = pd.read_csv(gt_csv)\n        sel = df[df[\"tomo_id\"] == selected_tomo_name]\n        # ✅ FIX: Load in 0, 1, 2 order to perfectly match numpy's Z, Y, X shape\n        if {\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"}.issubset(sel.columns):\n            gt_points = sel[[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values.astype(float)\n\n    if gt_points.shape[0] == 0:\n        st.warning(\"⚠️ No GT points found — metrics will be NaN.\")\n\n    combos = _make_grid(patch_sizes, stride_factors, dbscan_eps_list, presets, max_runs)\n    st.info(f\"🧪 Running {len(combos)} ablations on {selected_tomo_name}\")\n    prog = st.progress(0)\n    start = time.time()\n\n    for i, (ps, sf, eps, pr) in enumerate(combos):\n        if (time.time() - start) / 60 > max_total_runtime:\n            st.warning(\"⏱️ Global time limit reached.\")\n            break\n        try:\n            metrics = _run_single_combo(\n                volume_memmap, gt_points, ps, sf, eps, pr, model_choice, max_runtime_per_run * 60\n            )\n            metrics.update({\"tomo\": selected_tomo_name, \"run_index\": i + 1})\n            st.session_state[\"ablation_results\"].append(metrics)\n            st.write(\"✅\", metrics)\n        except Exception as e:\n            st.error(f\"Run failed: {e}\")\n            st.session_state[\"ablation_results\"].append({\n                \"tomo\": selected_tomo_name, \"run_index\": i + 1, \"patch_size\": ps,\n                \"stride_factor\": sf, \"dbscan_eps\": eps, \"preset\": pr, \"model\": model_choice,\n                \"error\": str(e), \"timestamp\": datetime.now().isoformat(),\n            })\n        prog.progress(int(100 * (i + 1) / len(combos)))\n\n    _save_results(st.session_state[\"ablation_results\"], \"ablation_final\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:11.541663Z","iopub.execute_input":"2026-02-22T14:41:11.542302Z","iopub.status.idle":"2026-02-22T14:41:11.548159Z","shell.execute_reply.started":"2026-02-22T14:41:11.542283Z","shell.execute_reply":"2026-02-22T14:41:11.547378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"NGROK_AUTH_TOKEN\"] = \"3A1jjzQstZSEnCZrTCNVTRDMts7_7rq2hNBqtWmmvS3nwGb9v\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:11.548765Z","iopub.execute_input":"2026-02-22T14:41:11.548950Z","iopub.status.idle":"2026-02-22T14:41:11.571163Z","shell.execute_reply.started":"2026-02-22T14:41:11.548935Z","shell.execute_reply":"2026-02-22T14:41:11.570438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ✅ Secure Streamlit + Ngrok Launcher (Kaggle Stable Version)\n# ============================================================\n# Run AFTER saving /kaggle/working/app.py\n\n# ---------------- Install Dependency ----------------\n!pip install -q pyngrok\n\n# ---------------- Imports ----------------\nimport os\nimport time\nimport torch\nimport platform\nimport subprocess\nimport threading\nfrom pathlib import Path\nfrom pyngrok import ngrok, conf\n\n# ---------------- Paths ----------------\nAPP_PATH = \"/kaggle/working/app.py\"\nLOG_FILE = \"/kaggle/working/streamlit_log.txt\"\n\nCACHE_DIR = Path(\"/kaggle/working/cache_phase1\")\nLOG_DIR   = CACHE_DIR / \"experiment_logs\"\nAUG_DIR   = CACHE_DIR / \"augmented\"\n\nfor d in [CACHE_DIR, LOG_DIR, AUG_DIR]:\n    d.mkdir(parents=True, exist_ok=True)\n\n# ---------------- Validate App Exists ----------------\nif not Path(APP_PATH).exists():\n    raise FileNotFoundError(f\"❌ {APP_PATH} not found. Save app.py first.\")\n\n# ---------------- Kill Old Sessions ----------------\nprint(\"🧹 Cleaning previous Streamlit/Ngrok processes...\")\nos.system(\"pkill -9 -f streamlit || true\")\nos.system(\"pkill -9 -f ngrok || true\")\ntime.sleep(2)\n\n# ---------------- Environment Summary ----------------\nprint(\"\\n🧠 Environment Summary\")\nprint(f\"Python: {platform.python_version()}\")\nprint(f\"Torch: {torch.__version__}\")\nprint(f\"CUDA Available: {torch.cuda.is_available()}\")\n\nif torch.cuda.is_available():\n    props = torch.cuda.get_device_properties(0)\n    print(f\"GPU: {props.name}\")\n    print(f\"VRAM: {round(props.total_memory / 1e9, 2)} GB\")\n\n# ---------------- Ngrok Authentication ----------------\nNGROK_AUTH_TOKEN = os.environ.get(\"NGROK_AUTH_TOKEN\")\n\nif not NGROK_AUTH_TOKEN:\n    raise ValueError(\n        \"❌ NGROK_AUTH_TOKEN not found.\\n\"\n        \"Run this first:\\n\"\n        'os.environ[\"NGROK_AUTH_TOKEN\"] = \"your_token_here\"'\n    )\n\nconf.get_default().auth_token = NGROK_AUTH_TOKEN\nprint(\"🔑 Ngrok authentication successful.\")\n\n# ---------------- Launch Streamlit ----------------\ndef run_streamlit():\n    with open(LOG_FILE, \"w\") as log:\n        subprocess.run(\n            [\n                \"streamlit\", \"run\", APP_PATH,\n                \"--server.headless\", \"true\",\n                \"--server.port\", \"8501\",\n                \"--server.address\", \"0.0.0.0\",\n                \"--browser.gatherUsageStats\", \"false\"\n            ],\n            stdout=log,\n            stderr=subprocess.STDOUT\n        )\n\nprint(\"\\n🚀 Starting Streamlit server...\")\nthreading.Thread(target=run_streamlit, daemon=True).start()\n\n# Give Streamlit time to initialize\ntime.sleep(8)\n\n# ---------------- Start Ngrok Tunnel ----------------\ntry:\n    tunnel = ngrok.connect(8501)\n    print(\"\\n🌐 Public URL:\")\n    print(\"👉\", tunnel.public_url)\nexcept Exception as e:\n    print(f\"❌ Ngrok failed: {e}\")\n\nprint(\"\\n✅ Launcher completed. Wait a few seconds if app is still loading.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:11.572889Z","iopub.execute_input":"2026-02-22T14:41:11.573110Z","iopub.status.idle":"2026-02-22T14:41:25.625497Z","shell.execute_reply.started":"2026-02-22T14:41:11.573087Z","shell.execute_reply":"2026-02-22T14:41:25.624668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python -m py_compile /kaggle/working/app.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:25.626716Z","iopub.execute_input":"2026-02-22T14:41:25.627435Z","iopub.status.idle":"2026-02-22T14:41:25.875777Z","shell.execute_reply.started":"2026-02-22T14:41:25.627409Z","shell.execute_reply":"2026-02-22T14:41:25.874690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cat /kaggle/working/streamlit_log.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T14:41:25.877202Z","iopub.execute_input":"2026-02-22T14:41:25.877447Z","iopub.status.idle":"2026-02-22T14:41:26.005444Z","shell.execute_reply.started":"2026-02-22T14:41:25.877420Z","shell.execute_reply":"2026-02-22T14:41:26.004741Z"}},"outputs":[],"execution_count":null}]}