{"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":"2025-10-31T14:12:29.856425Z","iopub.execute_input":"2025-10-31T14:12:29.856683Z","iopub.status.idle":"2025-10-31T14:12:35.307243Z","shell.execute_reply.started":"2025-10-31T14:12:29.856652Z","shell.execute_reply":"2025-10-31T14:12:35.306519Z"}},"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":"2025-10-31T14:12:35.308534Z","iopub.execute_input":"2025-10-31T14:12:35.308768Z"}},"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\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([np.array(Image.open(s)) for s in slices], axis=0).astype(np.float32)\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# ---------------------------------------------------------------\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 safe downsampling fallback.\"\"\"\n    if measure is None:\n        return None, None\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    vol = np.array(mm)\n    if level is None:\n        level = vol.mean() + 0.5 * vol.std()\n    try:\n        verts, faces, _, _ = measure.marching_cubes(vol, level=level)\n    except Exception as e:\n        warnings.warn(f\"marching_cubes failed: {e}\")\n        return None, None\n    if verts.shape[0] > target_max_vertices:\n        factor = int(np.ceil((verts.shape[0] / target_max_vertices) ** (1/3)))\n        if factor > 1:\n            vol_ds = vol[::factor, ::factor, ::factor]\n            try:\n                verts, faces, _, _ = measure.marching_cubes(vol_ds, level=level)\n                verts *= factor\n            except Exception as e:\n                warnings.warn(f\"Downsampled marching_cubes failed: {e}\")\n                idx = np.random.choice(verts.shape[0], size=target_max_vertices, replace=False)\n                verts = verts[idx]\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    \"\"\"\n    Convert voxel indices (x,y,z) → nanometers.\n    coords: Nx3 (x,y,z)\n    voxel_size_nm: (z_nm, y_nm, x_nm)\n    \"\"\"\n    if coords is None or coords.size == 0:\n        return coords\n    z_nm, y_nm, x_nm = voxel_size_nm\n    scale = np.array([x_nm, y_nm, z_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    inv = np.array([1/x_nm, 1/y_nm, 1/z_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# ================================================================\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    Z, Y, X = volume_arr.shape\n    patches_v, patches_t = [], []\n    for _ in range(n_patches):\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        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    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    from torch.cuda.amp import autocast\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                    with autocast():\n                        out = model(inp)\n                    out_np = out[0, 0].cpu().numpy()\n                    # 🧠 compute patch-level confidence (mean activation)\n                    conf = float(out_np.mean())\n                    if conf < 0.25:  # ⚖️ skip low-confidence patches → better recall balance\n                        continue\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                    del inp, out, out_np\n                    torch.cuda.empty_cache()\n    norm[norm == 0] = 1.0\n    conf_map /= norm\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\n# ================================================================\n# Part 3 – Configurable Ablation Controller (Streamlit Integrated, nm-aware, Stable)\n# ================================================================\n\nimport itertools, json, csv, 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 for Ablation Sweep ----------------\nst.sidebar.header(\"⚙️ Ablation Controller\")\nwith st.sidebar.expander(\"Sweep Settings\", expanded=True):\n    patch_sizes = st.multiselect(\"Patch sizes (voxels)\",\n                                 options=[32, 48, 64, 80], default=[48, 64])\n    stride_factors = st.multiselect(\"Stride factors (speed tradeoff)\",\n                                    options=[0.75, 1.0, 1.25, 1.5], default=[1.0, 1.25])\n    dbscan_eps_list = st.multiselect(\"DBSCAN eps (voxels)\",\n                                     options=[1.0, 1.5, 2.0, 3.0], default=[1.5, 2.0])\n    presets = st.multiselect(\"Training preset\",\n                             options=[\"fast\", \"balanced\", \"thorough\"], default=[\"balanced\"])\n    max_runs = st.number_input(\"Max total runs (safety cap)\", min_value=1, max_value=200, value=8, step=1)\n    max_runtime_per_run = st.slider(\"Max runtime per run (minutes)\", min_value=5, max_value=120, value=30, step=5)\n    max_total_runtime = st.number_input(\"Max total runtime (minutes)\", min_value=10, max_value=480, value=120)\n\nst.sidebar.markdown(\"---\")\nrun_button = st.sidebar.button(\"🚀 Run Ablation Sweep\")\nsave_button = st.sidebar.button(\"💾 Save Current Results (.csv/.json)\")\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, max_runs_allowed):\n    combos = list(itertools.product(patches, strides, eps_list, presets_list))\n    if len(combos) > max_runs_allowed:\n        combos = combos[:max_runs_allowed]\n    return combos\n\ndef _save_results_file(results_list, prefix=\"ablation_results\"):\n    \"\"\"Save results list to CSV + JSON.\"\"\"\n    if not results_list:\n        return None\n    df = pd.DataFrame(results_list)\n    ts = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n    csv_path = Path(CACHE_DIR) / f\"{prefix}_{ts}.csv\"\n    json_path = Path(CACHE_DIR) / f\"{prefix}_{ts}.json\"\n    df.to_csv(csv_path, index=False)\n    with open(json_path, \"w\") as f:\n        json.dump(results_list, f, indent=2)\n    return str(csv_path), str(json_path)\n\ndef _run_single_combo(volume_path_or_arr, gt_points,\n                      patch_size, stride_factor, dbscan_eps, preset, max_runtime_s):\n    \"\"\"Run one ablation configuration with time limits.\"\"\"\n    global DBSCAN_EPS, PATCH_SIZE, PATCH_STRIDE\n    prev_eps, prev_patch, prev_stride = DBSCAN_EPS, PATCH_SIZE, PATCH_STRIDE\n    DBSCAN_EPS, PATCH_SIZE, PATCH_STRIDE = float(dbscan_eps), int(patch_size), int(patch_size // 2)\n\n    try:\n        t0 = time.time()\n        clustered, metrics = run_tiny_cnn(\n            volume_path_or_arr,\n            gt_points,\n            epochs=EPOCHS,\n            preset=preset,\n            stride_factor=stride_factor\n        )\n        # nm-aware metric summary fields\n        metrics.update({\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        return clustered, metrics\n    finally:\n        DBSCAN_EPS, PATCH_SIZE, PATCH_STRIDE = prev_eps, prev_patch, prev_stride\n        gc.collect(); torch.cuda.empty_cache()\n\n# ---------------- Main Ablation Runner ----------------\nif run_button:\n    tomos = find_tomograms(DATA_ROOT)\n    if len(tomos) == 0:\n        st.error(\"❌ No tomograms found. Check DATA_ROOT.\")\n    else:\n        colA, colB = st.columns([1, 3])\n        with colA:\n            tomo_names = [t.name for t in tomos]\n            selected_tomo_name = st.selectbox(\"Select Tomogram for Ablation\", tomo_names, index=0)\n        selected_tomo = [t for t in tomos if t.name == selected_tomo_name][0]\n\n        # Load memmap + GT\n        try:\n            raw_memmap = load_volume_from_jpegs_cached(str(selected_tomo))\n            proc_cache_path = Path(CACHE_DIR) / f\"{selected_tomo.name}_proc.npy\"\n            if proc_cache_path.exists():\n                volume_memmap = np.load(proc_cache_path, mmap_mode=\"r\")\n            else:\n                volume_memmap = preprocess_volume_cached_memmap(raw_memmap)\n                proc_cache_path = Path(volume_memmap.filename)\n        except Exception as e:\n            st.error(f\"❌ Failed to load tomogram: {e}\")\n            volume_memmap = None\n\n        gt_points = np.zeros((0, 3))\n        gt_csv = Path(DATA_ROOT) / GT_CSV\n        if gt_csv.exists():\n            try:\n                df = pd.read_csv(gt_csv)\n                sel = df[df[\"tomo_id\"] == selected_tomo_name]\n                if {\"Motor axis 2\", \"Motor axis 1\", \"Motor axis 0\"}.issubset(sel.columns):\n                    gt_points = sel[[\"Motor axis 2\", \"Motor axis 1\", \"Motor axis 0\"]].values.astype(float)\n            except Exception:\n                gt_points = np.zeros((0, 3))\n\n        if volume_memmap is None:\n            st.error(\"⚠️ No volume loaded — aborting ablation run.\")\n        else:\n            combos = _make_grid(patch_sizes, stride_factors, dbscan_eps_list, presets, max_runs)\n            st.info(f\"🧪 Starting ablation sweep: {len(combos)} configurations on {selected_tomo_name}\")\n\n            progress_bar = st.progress(0)\n            results_local = []\n            sweep_start = time.time()\n\n            for idx, (ps, sf, eps, pr) in enumerate(combos):\n                elapsed_min = (time.time() - sweep_start) / 60.0\n                if elapsed_min > float(max_total_runtime):\n                    st.warning(\"⏱️ Global sweep time limit reached — stopping early.\")\n                    break\n\n                st.write(f\"Run {idx+1}/{len(combos)} → Patch={ps}, Stride={sf}, EPS={eps}, Preset={pr}\")\n                try:\n                    clustered, metrics = _run_single_combo(\n                        volume_memmap, gt_points, patch_size=ps,\n                        stride_factor=sf, dbscan_eps=eps,\n                        preset=pr, max_runtime_s=max_runtime_per_run * 60.0\n                    )\n                    metrics.update({\"tomo\": selected_tomo_name, \"run_index\": idx + 1})\n                    results_local.append(metrics)\n                    st.write(\"✅\", metrics)\n                except Exception as e:\n                    st.error(f\"Run failed for {ps, sf, eps, pr}: {e}\")\n                    results_local.append({\n                        \"tomo\": selected_tomo_name, \"run_index\": idx + 1,\n                        \"patch_size\": ps, \"stride_factor\": sf, \"dbscan_eps\": eps,\n                        \"preset\": pr, \"error\": str(e), \"timestamp\": datetime.now().isoformat()\n                    })\n\n                progress_bar.progress(int(100 * (idx + 1) / len(combos)))\n                st.session_state[\"ablation_results\"].extend(results_local)\n                _save_results_file(st.session_state[\"ablation_results\"], prefix=\"ablation_intermediate\")\n                results_local = []\n                gc.collect(); torch.cuda.empty_cache()\n\n            saved = _save_results_file(st.session_state[\"ablation_results\"], prefix=\"ablation_final\")\n            if saved:\n                st.success(f\"✅ Sweep complete. Results saved to {saved[0]} and {saved[1]}\")\n            else:\n                st.warning(\"⚠️ No results saved.\")\n\n# ---------------- Result Display ----------------\nst.header(\"📊 Ablation Results Summary (nm-normalized)\")\nif st.session_state.get(\"ablation_results\"):\n    df_results = pd.DataFrame(st.session_state[\"ablation_results\"])\n\n    # F₂ (nm-weighted) ranking for physical realism\n    if \"mean_dist_nm\" in df_results.columns:\n        df_results[\"weighted_f2\"] = df_results[\"f2\"] / (1 + df_results[\"mean_dist_nm\"].clip(lower=1))\n    else:\n        df_results[\"weighted_f2\"] = df_results[\"f2\"]\n\n    df_results = df_results.sort_values(by=[\"weighted_f2\"], ascending=False).reset_index(drop=True)\n    st.dataframe(df_results)\n\n    try:\n        topk = df_results.head(10)\n        fig_rank = px.bar(\n            topk,\n            x=topk.index.astype(str),\n            y=\"weighted_f2\",\n            hover_data=[\"precision\", \"recall\", \"mean_dist_nm\",\n                        \"patch_size\", \"preset\", \"stride_factor\"],\n            title=\"Top 10 Runs by Weighted F₂ (nm-normalized)\",\n            labels={\"x\": \"Run #\", \"weighted_f2\": \"Weighted F₂\"}\n        )\n        st.plotly_chart(fig_rank, use_container_width=True)\n    except Exception:\n        st.warning(\"⚠️ Could not render F₂ chart — check logged fields.\")\nelse:\n    st.info(\"No ablation results yet — run the sweep from the sidebar.\")\n\n# ---------------- Manual Save ----------------\nif save_button:\n    if st.session_state.get(\"ablation_results\"):\n        csvp, jsonp = _save_results_file(st.session_state[\"ablation_results\"], prefix=\"ablation_manual_save\")\n        st.success(f\"💾 Saved current ablation results to: {csvp} and {jsonp}\")\n    else:\n        st.warning(\"No results to save.\")\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\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            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        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                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(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 = sliding_window_inference(vol, model, patch_size=PATCH_SIZE, stride=int(PATCH_STRIDE * stride_factor))\n\n        # Threshold + clustering\n        thresh = recon.mean() + THRESH_STD * recon.std()\n        coords = np.array(np.nonzero(recon > thresh)).T\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        coords_xyz = coords[:, [2, 1, 0]].astype(float)\n        clustered = cluster_points(coords_xyz, eps=DBSCAN_EPS)\n\n        # Metrics: try nm-normalization if helper exists\n        metrics = compute_metrics(clustered, gt_points, radius=5.0)\n        # try to compute nm distance if API available\n        mean_dist_nm = None\n        try:\n            if \"get_default_voxel_size\" in globals():\n                vx = get_default_voxel_size(\"byu_motor\")  # user should implement / override\n                if vx is not None:\n                    mean_dist_nm = float(metrics.get(\"mean_dist\", np.nan)) * float(vx)\n                    metrics[\"mean_dist_nm\"] = mean_dist_nm\n        except Exception:\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\"\nif proc_cache_path.exists():\n    vol_memmap = np.load(proc_cache_path, mmap_mode=\"r\")\nelse:\n    vol_memmap = preprocess_volume_cached_memmap(raw_memmap)\n    proc_cache_path = Path(vol_memmap.filename)\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        if {\"Motor axis 2\", \"Motor axis 1\", \"Motor axis 0\"}.issubset(sel.columns):\n            gt_points = sel[[\"Motor axis 2\", \"Motor axis 1\", \"Motor axis 0\"]].values.astype(float)\n    except Exception:\n        pass\n\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.\")\n","metadata":{"trusted":true},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!grep -n \"global PATCH_SIZE\" /kaggle/working/app.py\n","metadata":{"trusted":true},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ✅ Secure Streamlit + Ngrok Launcher (7-Part Research Pipeline)\n# ============================================================\n# Run this cell in a Kaggle Notebook *after* saving /kaggle/working/app.py\n\n!pip install pyngrok --quiet\n\nimport os, time, threading, torch, platform, subprocess\nfrom pathlib import Path\nfrom pyngrok import ngrok, conf\n\n# ---------------- Directories ----------------\nAPP_PATH   = \"/kaggle/working/app.py\"\nCACHE_DIR  = Path(\"/kaggle/working/cache_phase1\")\nLOG_DIR    = CACHE_DIR / \"experiment_logs\"\nAUG_DIR    = CACHE_DIR / \"augmented\"\nfor d in [CACHE_DIR, LOG_DIR, AUG_DIR]:\n    d.mkdir(parents=True, exist_ok=True)\n\n# ---------------- Ngrok Authentication ----------------\nos.environ[\"NGROK_AUTH_TOKEN\"] = \"34nHhHaPtyV5ZSSSdLW4p2ugIVa_BuJDvKDo78ESrFi8jno5\"\n\nNGROK_AUTH_TOKEN = os.environ.get(\"NGROK_AUTH_TOKEN\", \"\").strip()\nif not NGROK_AUTH_TOKEN:\n    print(\"⚠️  No ngrok token detected — please run:\")\n    print('     os.environ[\"NGROK_AUTH_TOKEN\"] = \"your_token_here\"')\nelse:\n    conf.get_default().auth_token = NGROK_AUTH_TOKEN\n    print(\"🔑  Ngrok token loaded successfully.\")\n\n# ---------------- Clean up any old sessions ----------------\nos.system(\"pkill -f streamlit || true\")\ntry:\n    ngrok.kill()\nexcept Exception:\n    pass\n\n# ---------------- Environment Summary ----------------\nprint(\"🧠 Environment Summary:\")\nprint(f\"Python {platform.python_version()} | Torch {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    props = torch.cuda.get_device_properties(0)\n    print(f\"GPU: {props.name} — {round(props.total_memory / 1e9, 2)} GB VRAM × {torch.cuda.device_count()}\")\n\n# ---------------- Launch Streamlit via ngrok ----------------\ntry:\n    public_url = ngrok.connect(addr=8501)\n    print(f\"\\n🌐  Streamlit app public URL:\\n👉 {public_url}\")\nexcept Exception as e:\n    print(\"❌  Ngrok failed to connect:\", e)\n    print(\"Retrying in local-only mode...\")\n\ndef run_streamlit():\n    cmd = [\n        \"streamlit\", \"run\", APP_PATH,\n        \"--server.headless\", \"true\",\n        \"--server.port\", \"8501\",\n        \"--browser.gatherUsageStats\", \"false\",\n        \"--theme.base\", \"light\",\n    ]\n    subprocess.run(cmd)\n\nthread = threading.Thread(target=run_streamlit, daemon=True)\nthread.start()\n\ntime.sleep(8)\nprint(\"\\n⚙️  Streamlit is starting… this may take 10–20 seconds.\")\nprint(\"📁 Cache dir:\", CACHE_DIR)\nprint(\"📁 Logs dir :\", LOG_DIR)\nprint(\"📁 Aug dir  :\", AUG_DIR)\nprint(\"⚠️  Keep this cell running — closing it stops the public app.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}