| from __future__ import annotations |
|
|
| import contextlib |
| import io |
| import os |
| import re |
| import shutil |
| import stat |
| import tempfile |
| import time |
| from pathlib import Path |
|
|
| import nibabel as nib |
| import numpy as np |
| import streamlit as st |
|
|
| from data_prep.prep_utils import ( |
| DatasetConfig, |
| combine_standardized, |
| combine_standardized_images_only, |
| run_prep, |
| run_prep_images_only, |
| ) |
| from inference.streamlit_panel import render_inference_panel |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| WEBAPP_UPLOAD_ROOT = PROJECT_ROOT / "data" / "raw" / "webapp_upload" |
| PREP_OUT_ROOT = PROJECT_ROOT / "data" / "prep_outputs" |
| MODEL_INPUT_DEST = PROJECT_ROOT / "data" / "processed" / "test_input" |
| MODEL_OUTPUT_DEST = PROJECT_ROOT / "data" / "processed" / "test_output" |
| HF_RUNTIME_CACHE_ROOT = PROJECT_ROOT / "data" / "runtime_bundle_from_hf" |
| HF_RUNTIME_DEFAULT_REPO = os.environ.get("HF_RUNTIME_REPO_ID", "rmbielski/Stroke_Dependencies") |
| HF_RUNTIME_DEFAULT_REVISION = os.environ.get("HF_RUNTIME_REVISION", "main") |
| HF_MODEL_CACHE_ROOT = PROJECT_ROOT / "data" / "model_bundle_from_hf" |
| HF_MODEL_DEFAULT_REPO = os.environ.get("HF_MODEL_REPO_ID", "rmbielski/ARC_ATLAS_v3.1") |
| HF_MODEL_DEFAULT_REVISION = os.environ.get("HF_MODEL_REVISION", HF_RUNTIME_DEFAULT_REVISION) |
| HF_MODEL_DEFAULT_SUBDIR = os.environ.get("HF_MODEL_SUBDIR", "") |
| HF_DATA_CACHE_ROOT = PROJECT_ROOT / "data" / "hf_data_cache" |
| PUBLIC_HF_DATASETS = ["rmbielski/Atlas_2", "rmbielski/ARC", "rmbielski/Approximate_Numeracy"] |
| NATIVE_PREVIEW_ROOT = PROJECT_ROOT / "data" / "raw" / "webapp_native_preview" |
| ANTS_PROFILE_CHOICES = ["quick", "fast", "balanced", "accurate"] |
| HF_RUNTIME_ALLOW_PATTERNS = [ |
| "tools/ants/**", |
| "data/templateflow/**", |
| "**/tools/ants/**", |
| "**/data/templateflow/**", |
| ] |
|
|
|
|
| @st.cache_data(show_spinner=False) |
| def _load_nifti(path: str): |
| img = nib.load(path) |
| data = np.asarray(img.dataobj) |
| zooms = img.header.get_zooms() |
| return data, tuple(float(z) for z in zooms[:3]) |
|
|
|
|
| @st.cache_data(show_spinner=False) |
| def _load_nifti_from_bytes(file_name: str, payload: bytes): |
| suffix = ".nii.gz" if file_name.lower().endswith(".nii.gz") else ".nii" |
| with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: |
| tmp.write(payload) |
| tmp_path = tmp.name |
| try: |
| return _load_nifti(tmp_path) |
| finally: |
| try: |
| os.remove(tmp_path) |
| except OSError: |
| pass |
|
|
|
|
| def _is_nifti_filename(name: str) -> bool: |
| lower = name.lower() |
| return lower.endswith(".nii") or lower.endswith(".nii.gz") |
|
|
|
|
| def _clean_name(name: str) -> str: |
| cleaned = re.sub(r"[^a-zA-Z0-9._-]+", "_", name.strip()) |
| return cleaned or "input.nii.gz" |
|
|
|
|
| def _save_upload(uploaded_file, directory: Path) -> Path: |
| directory.mkdir(parents=True, exist_ok=True) |
| out_path = directory / _clean_name(uploaded_file.name) |
| out_path.write_bytes(uploaded_file.getbuffer()) |
| return out_path |
|
|
|
|
| def _clear_dir(path: Path) -> None: |
| if path.exists(): |
| shutil.rmtree(path) |
| path.mkdir(parents=True, exist_ok=True) |
|
|
|
|
| def _ensure_executable_binary(path: Path) -> Path: |
| if not path.exists(): |
| raise FileNotFoundError(f"Missing binary: {path}") |
|
|
| |
| try: |
| mode = path.stat().st_mode |
| path.chmod(mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH) |
| except Exception: |
| pass |
|
|
| if os.access(path, os.X_OK): |
| return path |
|
|
| |
| exec_dir = Path(tempfile.gettempdir()) / "stroke_runtime_exec" |
| exec_dir.mkdir(parents=True, exist_ok=True) |
| fallback = exec_dir / path.name |
| shutil.copy2(path, fallback) |
| fallback.chmod(fallback.stat().st_mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH) |
|
|
| if os.access(fallback, os.X_OK): |
| return fallback |
|
|
| raise PermissionError( |
| f"Binary is not executable: {path}. " |
| "This may be due to file permissions or a non-executable mount." |
| ) |
|
|
|
|
| def _persist_native_preview(uploaded_t1, uploaded_mask=None): |
| t1_dir = NATIVE_PREVIEW_ROOT / "t1" |
| mask_dir = NATIVE_PREVIEW_ROOT / "mask" |
| _clear_dir(t1_dir) |
| _save_upload(uploaded_t1, t1_dir) |
|
|
| if uploaded_mask is not None: |
| _clear_dir(mask_dir) |
| _save_upload(uploaded_mask, mask_dir) |
| elif mask_dir.exists(): |
| shutil.rmtree(mask_dir, ignore_errors=True) |
|
|
|
|
| def _load_persisted_native_preview(): |
| t1_dir = NATIVE_PREVIEW_ROOT / "t1" |
| mask_dir = NATIVE_PREVIEW_ROOT / "mask" |
| if not t1_dir.exists(): |
| return None, None, None |
|
|
| t1_paths = sorted(t1_dir.glob("*.nii*")) |
| if not t1_paths: |
| return None, None, None |
|
|
| t1_path = t1_paths[-1] |
| t1_data, t1_zooms = _load_nifti(str(t1_path)) |
|
|
| mask_data = None |
| mask_paths = sorted(mask_dir.glob("*.nii*")) if mask_dir.exists() else [] |
| if mask_paths: |
| mask_data, _ = _load_nifti(str(mask_paths[-1])) |
| return t1_data, t1_zooms, mask_data |
|
|
|
|
| def _discover_runtime_bundle_root(download_root: Path) -> Path: |
| direct_reg = download_root / "tools" / "ants" / "bin" / "antsRegistration" |
| if direct_reg.exists(): |
| return download_root |
|
|
| candidates = sorted(download_root.rglob("tools/ants/bin/antsRegistration")) |
| for reg in candidates: |
| bundle_root = reg.parents[3] |
| if (bundle_root / "tools" / "ants" / "bin" / "antsApplyTransforms").exists(): |
| return bundle_root |
| raise FileNotFoundError( |
| "Could not find runtime bundle layout under downloaded HF dataset. " |
| "Expected `tools/ants/bin/antsRegistration` and `tools/ants/bin/antsApplyTransforms`." |
| ) |
|
|
|
|
| def _set_runtime_env(bundle_root: Path) -> dict[str, str]: |
| reg = bundle_root / "tools" / "ants" / "bin" / "antsRegistration" |
| app = bundle_root / "tools" / "ants" / "bin" / "antsApplyTransforms" |
| tf_home = bundle_root / "data" / "templateflow" |
|
|
| missing = [str(p) for p in (reg, app, tf_home) if not p.exists()] |
| if missing: |
| raise FileNotFoundError( |
| "Runtime bundle is missing required paths:\n" + "\n".join(missing) |
| ) |
|
|
| reg_exec = _ensure_executable_binary(reg) |
| app_exec = _ensure_executable_binary(app) |
|
|
| os.environ["ANTS_REG"] = str(reg_exec) |
| os.environ["ANTS_APPLY"] = str(app_exec) |
| os.environ["TEMPLATEFLOW_HOME"] = str(tf_home) |
| |
| os.environ["ANTS_AUTO_INSTALL"] = "0" |
|
|
| return { |
| "runtime_bundle_root": str(bundle_root), |
| "ants_registration": str(reg_exec), |
| "ants_apply_transforms": str(app_exec), |
| "templateflow_home": str(tf_home), |
| } |
|
|
|
|
| def _hf_local_dir_for_repo(repo_id: str, revision: str | None = None) -> Path: |
| safe = re.sub(r"[^a-zA-Z0-9._-]+", "__", repo_id.strip()) |
| safe_rev = re.sub(r"[^a-zA-Z0-9._-]+", "__", (revision or "main").strip() or "main") |
| return HF_RUNTIME_CACHE_ROOT / safe / safe_rev |
|
|
|
|
| class _HFFileWrapper: |
| """Minimal shim that satisfies the UploadedFile interface (.name + .getbuffer()).""" |
|
|
| def __init__(self, name: str, data: bytes) -> None: |
| self.name = name |
| self._data = data |
|
|
| def getbuffer(self) -> bytes: |
| return self._data |
|
|
|
|
| def _list_nifti_files_from_hf( |
| repo_id: str, |
| revision: str = "main", |
| token: str | None = None, |
| ) -> tuple[list[str], list[str]]: |
| """Return (t1_files, mask_files) lists of relative NIfTI paths in the dataset.""" |
| try: |
| from huggingface_hub import list_repo_files |
| except Exception as exc: |
| raise RuntimeError( |
| "huggingface_hub is required. Install with `pip install huggingface_hub`." |
| ) from exc |
|
|
| all_files = list( |
| list_repo_files( |
| repo_id=repo_id.strip(), |
| repo_type="dataset", |
| revision=revision or "main", |
| token=(token.strip() if token else None), |
| ) |
| ) |
|
|
| nifti = sorted( |
| f for f in all_files if f.lower().endswith(".nii") or f.lower().endswith(".nii.gz") |
| ) |
|
|
| _MASK_KEYWORDS = {"mask", "lesion", "seg", "label"} |
| mask_files = [f for f in nifti if any(k in f.lower() for k in _MASK_KEYWORDS)] |
| t1_files = [f for f in nifti if f not in set(mask_files)] |
| return t1_files, mask_files |
|
|
|
|
| def _download_hf_nifti( |
| repo_id: str, |
| filename: str, |
| revision: str = "main", |
| token: str | None = None, |
| ) -> bytes: |
| """Download a single NIfTI file from an HF dataset and return its bytes.""" |
| try: |
| from huggingface_hub import hf_hub_download |
| except Exception as exc: |
| raise RuntimeError( |
| "huggingface_hub is required. Install with `pip install huggingface_hub`." |
| ) from exc |
|
|
| safe_repo = re.sub(r"[^a-zA-Z0-9._-]+", "__", repo_id.strip()) |
| cache_dir = HF_DATA_CACHE_ROOT / safe_repo / (revision or "main") |
| cache_dir.mkdir(parents=True, exist_ok=True) |
|
|
| local_path = hf_hub_download( |
| repo_id=repo_id.strip(), |
| filename=filename, |
| repo_type="dataset", |
| revision=revision or "main", |
| token=(token.strip() if token else None), |
| local_dir=str(cache_dir), |
| ) |
| return Path(local_path).read_bytes() |
|
|
|
|
| def _sync_runtime_bundle_from_hf( |
| repo_id: str, |
| revision: str, |
| token: str | None = None, |
| force_download: bool = False, |
| ) -> dict[str, str]: |
| try: |
| from huggingface_hub import snapshot_download |
| except Exception as exc: |
| raise RuntimeError( |
| "huggingface_hub is required to download runtime dependencies. " |
| "Install it with `pip install huggingface_hub`." |
| ) from exc |
|
|
| repo_id = repo_id.strip() |
| if not repo_id: |
| raise ValueError("Runtime dataset repo_id is empty.") |
| revision = (revision or "main").strip() or "main" |
|
|
| local_dir = _hf_local_dir_for_repo(repo_id, revision) |
| if force_download and local_dir.exists(): |
| shutil.rmtree(local_dir, ignore_errors=True) |
| local_dir.mkdir(parents=True, exist_ok=True) |
|
|
| if not force_download: |
| try: |
| bundle_root = _discover_runtime_bundle_root(local_dir) |
| except FileNotFoundError: |
| pass |
| else: |
| info = _set_runtime_env(bundle_root) |
| info["hf_repo_id"] = repo_id |
| info["hf_revision"] = revision |
| info["hf_local_dir"] = str(local_dir) |
| info["hf_cache_hit"] = "true" |
| return info |
|
|
| snapshot_download( |
| repo_id=repo_id, |
| repo_type="dataset", |
| local_dir=str(local_dir), |
| allow_patterns=HF_RUNTIME_ALLOW_PATTERNS, |
| revision=revision, |
| token=(token.strip() if token else None), |
| ) |
| bundle_root = _discover_runtime_bundle_root(local_dir) |
| info = _set_runtime_env(bundle_root) |
| info["hf_repo_id"] = repo_id |
| info["hf_revision"] = revision |
| info["hf_local_dir"] = str(local_dir) |
| info["hf_cache_hit"] = "false" |
| return info |
|
|
|
|
| def _paired_mask_path(t1_path: Path, mask_dir: Path) -> Path | None: |
| if not mask_dir.exists(): |
| return None |
| key = t1_path.name.replace("_T1w_MNI_norm", "") |
| for mask_path in sorted(mask_dir.glob("*.nii*")): |
| mask_key = mask_path.name.replace("_lesion_mask_MNI_clean", "") |
| if mask_key == key: |
| return mask_path |
| return None |
|
|
|
|
| def _normalize_slice(slice_data: np.ndarray) -> np.ndarray: |
| low, high = np.percentile(slice_data, [1, 99]) |
| if low == high: |
| low, high = float(slice_data.min()), float(slice_data.max()) |
| if low == high: |
| low, high = 0.0, 1.0 |
| normalized = (slice_data - low) / (high - low) |
| return np.clip(normalized, 0.0, 1.0) |
|
|
|
|
| def _select_volume(data: np.ndarray, key_prefix: str, label: str) -> np.ndarray: |
| if data.ndim <= 3: |
| return data |
| if data.shape[-1] <= 1: |
| return np.squeeze(data, axis=-1) |
| vol_idx = st.slider( |
| f"{label} volume index", |
| 0, |
| data.shape[-1] - 1, |
| 0, |
| key=f"{key_prefix}_vol_idx", |
| ) |
| return data[..., vol_idx] |
|
|
|
|
| def _overlay_mask(normalized_slice: np.ndarray, mask_slice: np.ndarray, opacity: float) -> np.ndarray: |
| rgb = np.stack([normalized_slice, normalized_slice, normalized_slice], axis=-1) |
| mask_bool = np.asarray(mask_slice) > 0 |
| overlay_color = np.array([1.0, 0.0, 0.0], dtype=rgb.dtype) |
| rgb[mask_bool] = (1.0 - opacity) * rgb[mask_bool] + opacity * overlay_color |
| return rgb |
|
|
|
|
| def _render_inline_image(container, preview: np.ndarray, caption: str) -> None: |
| arr = np.asarray(preview, dtype=np.float32) |
| arr = np.clip(np.nan_to_num(arr, nan=0.0, posinf=1.0, neginf=0.0), 0.0, 1.0) |
| container.image((arr * 255.0).astype(np.uint8), caption=caption, use_container_width=True) |
|
|
|
|
| def _render_scaled_image(preview: np.ndarray, caption: str, key_prefix: str): |
| left, _ = st.columns([1, 1]) |
| _render_inline_image(left, preview, caption) |
|
|
|
|
| def _render_overlay_from_arrays( |
| image_data: np.ndarray, |
| image_zooms: tuple[float, float, float], |
| mask_data: np.ndarray | None, |
| key_prefix: str, |
| title: str, |
| ): |
| st.markdown(f"**{title}**") |
|
|
| image_vol = _select_volume(image_data, key_prefix, title) |
| if image_vol.ndim < 2: |
| st.error(f"Unsupported image dimensions: {image_data.shape}") |
| return |
|
|
| col_a, col_b, col_c = st.columns(3) |
| col_a.metric("Image shape", str(image_data.shape)) |
| col_b.metric("Image dtype", str(image_data.dtype)) |
| col_c.metric("Voxel size", " x ".join(f"{z:.3g}" for z in image_zooms)) |
|
|
| if image_vol.ndim == 2: |
| slice_data = np.asarray(image_vol) |
| caption = "2D image" |
| axis = None |
| idx = None |
| else: |
| axis_options = {"Axial (Z)": 2, "Coronal (Y)": 1, "Sagittal (X)": 0} |
| axis_label = st.selectbox( |
| f"{title} slice axis", |
| list(axis_options.keys()), |
| key=f"{key_prefix}_axis", |
| ) |
| axis = axis_options[axis_label] |
| idx = st.slider( |
| f"{title} slice index", |
| 0, |
| image_vol.shape[axis] - 1, |
| image_vol.shape[axis] // 2, |
| key=f"{key_prefix}_slice", |
| ) |
| slice_data = np.take(image_vol, idx, axis=axis) |
| caption = f"{axis_label} slice {idx}" |
|
|
| preview = _normalize_slice(np.asarray(slice_data)) |
|
|
| if mask_data is not None: |
| mask_vol = _select_volume(mask_data, f"{key_prefix}_mask", f"{title} mask") |
| mask_slice = None |
| if axis is None: |
| if mask_vol.shape == image_vol.shape: |
| mask_slice = mask_vol |
| else: |
| st.warning("Mask shape does not match image shape. Showing image only.") |
| else: |
| if mask_vol.shape == image_vol.shape: |
| mask_slice = np.take(mask_vol, idx, axis=axis) |
| else: |
| st.warning( |
| f"Mask shape {mask_vol.shape} does not match image shape {image_vol.shape}. " |
| "Showing image only." |
| ) |
|
|
| if mask_slice is not None: |
| opacity = st.slider("Mask opacity", 0.0, 1.0, 0.4, 0.05, key=f"{key_prefix}_mask_alpha") |
| preview = _overlay_mask(preview, np.asarray(mask_slice), opacity) |
|
|
| _render_scaled_image(preview, caption=caption, key_prefix=key_prefix) |
|
|
|
|
| def _coerce_3d_for_compare(volume: np.ndarray) -> np.ndarray | None: |
| if volume.ndim == 3: |
| return volume |
| if volume.ndim == 4 and volume.shape[-1] == 1: |
| return np.squeeze(volume, axis=-1) |
| if volume.ndim >= 4: |
| return volume[..., 0] |
| return None |
|
|
|
|
| def _render_side_by_side_comparison( |
| native_data: np.ndarray, |
| native_mask: np.ndarray | None, |
| prepared_data: np.ndarray, |
| prepared_mask: np.ndarray | None, |
| ): |
| native_vol = _coerce_3d_for_compare(native_data) |
| prepared_vol = _coerce_3d_for_compare(prepared_data) |
|
|
| if native_vol is None or prepared_vol is None: |
| st.info( |
| "Side-by-side comparison currently supports 3D volumes (or 4D with a singleton/first volume)." |
| ) |
| return |
|
|
| st.markdown("**Native vs Registered (side-by-side)**") |
| axis_options = {"Axial (Z)": 2, "Coronal (Y)": 1, "Sagittal (X)": 0} |
| axis_label = st.selectbox("Comparison slice axis", list(axis_options.keys()), key="cmp_axis") |
| axis = axis_options[axis_label] |
|
|
| slice_pct = st.slider("Comparison slice position (%)", 0, 100, 50, key="cmp_slice_pct") |
| idx_native = int(round((native_vol.shape[axis] - 1) * (slice_pct / 100.0))) |
| idx_prepared = int(round((prepared_vol.shape[axis] - 1) * (slice_pct / 100.0))) |
|
|
| show_masks = st.checkbox("Overlay masks in comparison", value=True, key="cmp_show_masks") |
| opacity = st.slider("Comparison mask opacity", 0.0, 1.0, 0.4, 0.05, key="cmp_mask_alpha") |
|
|
| native_slice = np.take(native_vol, idx_native, axis=axis) |
| prepared_slice = np.take(prepared_vol, idx_prepared, axis=axis) |
| native_preview = _normalize_slice(np.asarray(native_slice)) |
| prepared_preview = _normalize_slice(np.asarray(prepared_slice)) |
|
|
| if show_masks and native_mask is not None: |
| native_mask_vol = _coerce_3d_for_compare(native_mask) |
| if native_mask_vol is not None and native_mask_vol.shape == native_vol.shape: |
| native_mask_slice = np.take(native_mask_vol, idx_native, axis=axis) |
| native_preview = _overlay_mask(native_preview, np.asarray(native_mask_slice), opacity) |
|
|
| if show_masks and prepared_mask is not None: |
| prepared_mask_vol = _coerce_3d_for_compare(prepared_mask) |
| if prepared_mask_vol is not None and prepared_mask_vol.shape == prepared_vol.shape: |
| prepared_mask_slice = np.take(prepared_mask_vol, idx_prepared, axis=axis) |
| prepared_preview = _overlay_mask(prepared_preview, np.asarray(prepared_mask_slice), opacity) |
|
|
| left, right = st.columns(2) |
| _render_inline_image( |
| left, |
| native_preview, |
| caption=f"Native {axis_label} slice {idx_native} / {native_vol.shape[axis] - 1}", |
| ) |
| _render_inline_image( |
| right, |
| prepared_preview, |
| caption=f"Registered {axis_label} slice {idx_prepared} / {prepared_vol.shape[axis] - 1}", |
| ) |
|
|
|
|
| def _run_registration(uploaded_t1, uploaded_mask, already_mni: bool, prep_profile: str) -> tuple[bool, str]: |
| t0 = time.monotonic() |
| _clear_dir(WEBAPP_UPLOAD_ROOT) |
| image_dir = WEBAPP_UPLOAD_ROOT / "Images" |
| mask_dir = WEBAPP_UPLOAD_ROOT / "Masks" |
|
|
| t1_saved = _save_upload(uploaded_t1, image_dir) |
| has_mask = uploaded_mask is not None |
| if has_mask: |
| _save_upload(uploaded_mask, mask_dir) |
|
|
| _clear_dir(MODEL_INPUT_DEST) |
|
|
| log_buffer = io.StringIO() |
| with contextlib.redirect_stdout(log_buffer): |
| os.environ["PREP_ANTS_PROFILE"] = (prep_profile or "quick").strip().lower() |
| print("[prep] staging complete") |
| print(f"[prep] uploaded_t1={t1_saved.name}") |
| print(f"[prep] uploaded_mask={'yes' if has_mask else 'no'}") |
| print(f"[prep] already_mni={already_mni}") |
| print(f"[prep] ants_profile={os.environ.get('PREP_ANTS_PROFILE')}") |
| print(f"[prep] command_timeout_sec={os.environ.get('PREP_CMD_TIMEOUT_SEC', os.environ.get('ANTS_CMD_TIMEOUT_SEC', '900'))}") |
| if has_mask: |
| datasets = [ |
| DatasetConfig( |
| name="WEBAPP", |
| images_dir=image_dir, |
| masks_dir=mask_dir, |
| t1_glob="**/*.nii*", |
| mask_glob="**/*.nii*", |
| already_mni=already_mni, |
| overwrite=True, |
| ) |
| ] |
| print("[prep] running image+mask prep pipeline...") |
| outputs = run_prep(datasets, PREP_OUT_ROOT, force_overwrite=True) |
| if not outputs: |
| raise RuntimeError("No outputs were generated from image+mask prep.") |
| print("[prep] combining standardized image+mask outputs...") |
| combine_standardized(outputs, MODEL_INPUT_DEST) |
| else: |
| print("[prep] running image-only prep pipeline...") |
| out_ds = run_prep_images_only( |
| images_dir=image_dir, |
| out_root=PREP_OUT_ROOT, |
| name="WEBAPP", |
| t1_glob="**/*.nii*", |
| already_mni=already_mni, |
| overwrite=True, |
| ) |
| print("[prep] combining standardized image-only outputs...") |
| combine_standardized_images_only([out_ds], MODEL_INPUT_DEST) |
|
|
| prepared_t1 = sorted((MODEL_INPUT_DEST / "t1").glob("*.nii*")) |
| if not prepared_t1: |
| raise RuntimeError( |
| "Prep finished but no registered/normalized images were found in " |
| f"{MODEL_INPUT_DEST / 't1'}." |
| ) |
|
|
| message = ( |
| f"Prepared {len(prepared_t1)} image(s). " |
| f"Model input root: {MODEL_INPUT_DEST}. " |
| f"Uploaded source: {t1_saved.name}" |
| ) |
| elapsed = time.monotonic() - t0 |
| return has_mask, message + f"\nPrep wall time: {elapsed:.2f}s\n\n" + log_buffer.getvalue() |
|
|
|
|
| def _render_hf_dataset_browser() -> None: |
| """Render the HF dataset browser inside the 'Browse HF Dataset' tab.""" |
| token = (st.session_state.get("hf_runtime_token") or "").strip() or None |
|
|
| dataset_options = PUBLIC_HF_DATASETS + ["Custom..."] |
| selected_preset = st.selectbox("Dataset", dataset_options, key="hf_browser_preset") |
|
|
| if selected_preset == "Custom...": |
| repo_id = st.text_input( |
| "Repository ID", |
| key="hf_browser_custom_repo", |
| placeholder="org/dataset-name", |
| ) |
| else: |
| repo_id = selected_preset |
|
|
| revision = st.text_input( |
| "Revision (branch / tag / commit)", |
| value="main", |
| key="hf_browser_revision", |
| ) |
|
|
| if token: |
| st.caption("HF token from Runtime Setup will be used (supports private repositories).") |
| else: |
| st.caption("For private repositories, enter your HF token in the Runtime Setup section above.") |
|
|
| if st.button("List files", key="hf_browser_list_btn") and repo_id: |
| with st.spinner(f"Listing NIfTI files in {repo_id}..."): |
| try: |
| t1_files, mask_files = _list_nifti_files_from_hf(repo_id, revision, token) |
| st.session_state["hf_browser_t1_files"] = t1_files |
| st.session_state["hf_browser_mask_files"] = mask_files |
| st.session_state["hf_browser_list_error"] = "" |
| st.session_state["hf_browser_repo_id"] = repo_id |
| st.session_state["hf_browser_revision_used"] = revision |
| except Exception as exc: |
| st.session_state["hf_browser_list_error"] = str(exc) |
| st.session_state["hf_browser_t1_files"] = [] |
| st.session_state["hf_browser_mask_files"] = [] |
|
|
| if st.session_state.get("hf_browser_list_error"): |
| st.error(st.session_state["hf_browser_list_error"]) |
|
|
| t1_files: list[str] = st.session_state.get("hf_browser_t1_files", []) |
| mask_files: list[str] = st.session_state.get("hf_browser_mask_files", []) |
|
|
| if t1_files or mask_files: |
| if not t1_files: |
| st.warning("No T1-like NIfTI files found. All NIfTI files were classified as masks.") |
| selected_t1 = None |
| else: |
| selected_t1 = st.selectbox("T1 image file", t1_files, key="hf_browser_selected_t1") |
|
|
| mask_options = ["— none —"] + mask_files |
| selected_mask_label = st.selectbox( |
| "Mask file (optional)", mask_options, key="hf_browser_selected_mask" |
| ) |
| selected_mask = None if selected_mask_label == "— none —" else selected_mask_label |
|
|
| if st.button("Load selected", key="hf_browser_load_btn", disabled=selected_t1 is None): |
| cached_repo = st.session_state.get("hf_browser_repo_id", repo_id) |
| cached_rev = st.session_state.get("hf_browser_revision_used", revision) |
| try: |
| with st.spinner(f"Downloading {Path(selected_t1).name}..."): |
| t1_bytes = _download_hf_nifti(cached_repo, selected_t1, cached_rev, token) |
| st.session_state["hf_loaded_t1"] = {"name": Path(selected_t1).name, "bytes": t1_bytes} |
| if selected_mask: |
| with st.spinner(f"Downloading {Path(selected_mask).name}..."): |
| mask_bytes = _download_hf_nifti(cached_repo, selected_mask, cached_rev, token) |
| st.session_state["hf_loaded_mask"] = { |
| "name": Path(selected_mask).name, |
| "bytes": mask_bytes, |
| } |
| else: |
| st.session_state["hf_loaded_mask"] = None |
| st.session_state["hf_load_error"] = "" |
| except Exception as exc: |
| st.session_state["hf_load_error"] = str(exc) |
|
|
| if st.session_state.get("hf_load_error"): |
| st.error(st.session_state["hf_load_error"]) |
|
|
| loaded_t1 = st.session_state.get("hf_loaded_t1") |
| loaded_mask = st.session_state.get("hf_loaded_mask") |
| if loaded_t1: |
| mask_label = f" | Mask: {loaded_mask['name']}" if loaded_mask else "" |
| st.success(f"Loaded: {loaded_t1['name']}{mask_label}") |
| if st.button("Clear HF selection", key="hf_browser_clear_btn"): |
| for k in ("hf_loaded_t1", "hf_loaded_mask", "hf_browser_t1_files", |
| "hf_browser_mask_files", "hf_load_error"): |
| st.session_state.pop(k, None) |
| st.rerun() |
|
|
|
|
| def render_app(): |
| st.set_page_config(page_title="Chronic Stroke Segmentation Prep", layout="wide") |
| st.title("Input Registration + Prep") |
| st.write( |
| "Upload a T1 NIfTI and optionally a lesion mask, preview native-space data, " |
| "run registration/normalization via `src/data_prep`, and compare outputs side-by-side." |
| ) |
|
|
| with st.expander("0) Runtime setup (Hugging Face dataset)", expanded=False): |
| st.write( |
| "Runtime dependencies are pulled from your HF dataset and cached locally. " |
| "These paths are used for every registration run." |
| ) |
| runtime_repo_id = st.text_input( |
| "Runtime dataset repo_id", |
| value=HF_RUNTIME_DEFAULT_REPO, |
| key="hf_runtime_repo_id", |
| help="Example: rmbielski/Stroke_Dependencies", |
| ) |
| runtime_revision = st.text_input( |
| "Dataset revision", |
| value=HF_RUNTIME_DEFAULT_REVISION, |
| key="hf_runtime_revision", |
| help="Branch, tag, or commit; usually `main`.", |
| ) |
| runtime_token = st.text_input( |
| "HF token (optional)", |
| value=os.environ.get("HF_TOKEN", ""), |
| type="password", |
| key="hf_runtime_token", |
| help="Required only for private datasets.", |
| ) |
| runtime_force = st.checkbox( |
| "Force re-download runtime bundle", |
| value=False, |
| key="hf_runtime_force", |
| ) |
| setup_btn = st.button("Sync runtime bundle from HF", key="setup_runtime_btn") |
| if setup_btn: |
| with st.spinner("Syncing runtime bundle from Hugging Face..."): |
| try: |
| runtime_info = _sync_runtime_bundle_from_hf( |
| repo_id=runtime_repo_id, |
| revision=runtime_revision, |
| token=runtime_token, |
| force_download=runtime_force, |
| ) |
| except Exception as exc: |
| st.session_state["runtime_setup_error"] = str(exc) |
| st.session_state["runtime_setup_info"] = None |
| else: |
| st.session_state["runtime_setup_error"] = "" |
| st.session_state["runtime_setup_info"] = runtime_info |
|
|
| if st.session_state.get("runtime_setup_error"): |
| st.error(st.session_state["runtime_setup_error"]) |
| if st.session_state.get("runtime_setup_info"): |
| st.success("Runtime bundle is available and active.") |
| st.code("\n".join(f"{k}: {v}" for k, v in st.session_state["runtime_setup_info"].items())) |
|
|
| if "native_upload_key_version" not in st.session_state: |
| st.session_state["native_upload_key_version"] = 0 |
| upload_key_version = int(st.session_state["native_upload_key_version"]) |
|
|
| with st.expander("1) Upload inputs", expanded=True): |
| upload_tab, hf_tab = st.tabs(["Upload file", "Browse HF Dataset"]) |
|
|
| with upload_tab: |
| uploaded_t1 = st.file_uploader( |
| "T1 MRI (.nii/.nii.gz)", |
| type=None, |
| key=f"prep_t1_{upload_key_version}", |
| ) |
| include_mask = st.checkbox("I also want to upload a mask", value=False, key="prep_has_mask") |
| uploaded_mask = st.file_uploader( |
| "Lesion mask (.nii/.nii.gz)", |
| type=None, |
| key=f"prep_mask_{upload_key_version}", |
| disabled=not include_mask, |
| ) |
|
|
| with hf_tab: |
| _render_hf_dataset_browser() |
|
|
| already_mni = st.checkbox( |
| "Input is already in MNI space (skip ANTs registration)", |
| value=False, |
| key="prep_already_mni", |
| ) |
| prep_profile = st.selectbox( |
| "Registration profile", |
| ANTS_PROFILE_CHOICES, |
| index=ANTS_PROFILE_CHOICES.index( |
| (os.environ.get("PREP_ANTS_PROFILE") or "quick").strip().lower() |
| if (os.environ.get("PREP_ANTS_PROFILE") or "quick").strip().lower() in ANTS_PROFILE_CHOICES |
| else "quick" |
| ), |
| key="prep_ants_profile", |
| ) |
| run_prep_btn = st.button("Register + normalize", type="primary", key="prep_run_btn") |
|
|
| |
| effective_t1 = uploaded_t1 |
| effective_mask = uploaded_mask if include_mask else None |
|
|
| if effective_t1 is None and st.session_state.get("hf_loaded_t1"): |
| d = st.session_state["hf_loaded_t1"] |
| effective_t1 = _HFFileWrapper(d["name"], d["bytes"]) |
|
|
| if effective_mask is None and st.session_state.get("hf_loaded_mask"): |
| d = st.session_state["hf_loaded_mask"] |
| effective_mask = _HFFileWrapper(d["name"], d["bytes"]) |
| include_mask = True |
|
|
| native_data, native_zooms, native_mask_data = _load_persisted_native_preview() |
|
|
| if effective_t1 is not None: |
| if not _is_nifti_filename(effective_t1.name): |
| st.error("T1 file must be `.nii` or `.nii.gz`.") |
| elif include_mask and effective_mask is not None and not _is_nifti_filename(effective_mask.name): |
| st.error("Mask file must be `.nii` or `.nii.gz`.") |
| else: |
| try: |
| _persist_native_preview( |
| uploaded_t1=effective_t1, |
| uploaded_mask=effective_mask, |
| ) |
| native_data, native_zooms, native_mask_data = _load_persisted_native_preview() |
| except Exception as exc: |
| st.error(f"Failed to persist native preview files: {exc}") |
|
|
| st.subheader("2) Native-space viewer (before registration)") |
| if st.button("Clear native viewer results", key="clear_native_results_btn"): |
| if NATIVE_PREVIEW_ROOT.exists(): |
| shutil.rmtree(NATIVE_PREVIEW_ROOT, ignore_errors=True) |
| st.session_state["native_upload_key_version"] = upload_key_version + 1 |
| st.rerun() |
|
|
| if native_data is None: |
| st.info("Upload a T1 file or load one from the HF Dataset browser above to preview the native-space MRI.") |
| else: |
| _render_overlay_from_arrays( |
| image_data=native_data, |
| image_zooms=native_zooms, |
| mask_data=native_mask_data, |
| key_prefix="native_view", |
| title="Native uploaded image", |
| ) |
|
|
| if run_prep_btn: |
| if effective_t1 is None: |
| st.error("Upload a T1 NIfTI or load one from the HF Dataset browser before running registration.") |
| elif not _is_nifti_filename(effective_t1.name): |
| st.error("T1 file must be `.nii` or `.nii.gz`.") |
| elif include_mask and effective_mask is None: |
| st.error("Mask is enabled but no mask file was provided.") |
| elif include_mask and effective_mask is not None and not _is_nifti_filename(effective_mask.name): |
| st.error("Mask file must be `.nii` or `.nii.gz`.") |
| else: |
| stage_status = st.empty() |
| t0_total = time.monotonic() |
| with st.spinner("Syncing runtime bundle + running registration pipeline..."): |
| try: |
| stage_status.info("Step 1/2: Resolving runtime bundle...") |
| t0_sync = time.monotonic() |
| runtime_info = _sync_runtime_bundle_from_hf( |
| repo_id=runtime_repo_id, |
| revision=runtime_revision, |
| token=runtime_token, |
| force_download=runtime_force, |
| ) |
| sync_elapsed = time.monotonic() - t0_sync |
| st.session_state["runtime_setup_error"] = "" |
| st.session_state["runtime_setup_info"] = runtime_info |
| cache_hit = str(runtime_info.get("hf_cache_hit", "false")).lower() == "true" |
| cache_text = "cache hit" if cache_hit else "downloaded" |
|
|
| stage_status.info( |
| "Step 2/2: Running registration + normalization " |
| f"(runtime {cache_text}, sync {sync_elapsed:.1f}s)..." |
| ) |
| t0_prep = time.monotonic() |
| has_mask, prep_log = _run_registration( |
| uploaded_t1=effective_t1, |
| uploaded_mask=effective_mask, |
| already_mni=already_mni, |
| prep_profile=prep_profile, |
| ) |
| prep_elapsed = time.monotonic() - t0_prep |
| except Exception as exc: |
| st.session_state["prep_error"] = str(exc) |
| st.session_state["prep_log"] = "" |
| st.session_state["prep_ready"] = False |
| st.session_state["prep_message"] = "" |
| stage_status.error("Registration pipeline failed. Check 'Prep logs' and traceback details.") |
| else: |
| total_elapsed = time.monotonic() - t0_total |
| st.session_state["prep_error"] = "" |
| runtime_summary = ( |
| "[runtime]\n" |
| f"hf_repo_id={runtime_info.get('hf_repo_id')}\n" |
| f"hf_revision={runtime_info.get('hf_revision')}\n" |
| f"hf_local_dir={runtime_info.get('hf_local_dir')}\n" |
| f"hf_cache_hit={runtime_info.get('hf_cache_hit')}\n" |
| f"already_mni={already_mni}\n" |
| f"ants_profile={os.environ.get('PREP_ANTS_PROFILE', 'balanced')}\n" |
| f"ants_registration={runtime_info.get('ants_registration')}\n" |
| f"ants_apply_transforms={runtime_info.get('ants_apply_transforms')}\n" |
| f"templateflow_home={runtime_info.get('templateflow_home')}\n" |
| f"prep_cmd_timeout_sec={os.environ.get('PREP_CMD_TIMEOUT_SEC', os.environ.get('ANTS_CMD_TIMEOUT_SEC', '900'))}\n" |
| "\n[timing]\n" |
| f"sync_seconds={sync_elapsed:.2f}\n" |
| f"prep_seconds={prep_elapsed:.2f}\n" |
| f"total_seconds={total_elapsed:.2f}\n" |
| ) |
| st.session_state["prep_log"] = runtime_summary + "\n" + prep_log |
| st.session_state["prep_ready"] = True |
| st.session_state["prep_has_mask_out"] = has_mask |
| st.session_state["prep_message"] = ( |
| "Registration/prep completed. Viewer and model-ready files were updated. " |
| f"Total time: {total_elapsed / 60.0:.1f} min." |
| ) |
| stage_status.success( |
| f"Registration pipeline completed in {total_elapsed / 60.0:.1f} min." |
| ) |
|
|
| if st.session_state.get("prep_error"): |
| st.error(st.session_state["prep_error"]) |
|
|
| if st.session_state.get("prep_message") and not st.session_state.get("prep_error"): |
| st.success(st.session_state["prep_message"]) |
|
|
| if st.session_state.get("prep_log"): |
| with st.expander("Prep logs", expanded=False): |
| st.text(st.session_state["prep_log"]) |
|
|
| st.subheader("3) Registered/normalized viewer") |
| if st.button("Clear registered results", key="clear_registered_results_btn"): |
| _clear_dir(MODEL_INPUT_DEST) |
| for k in ("prep_log", "prep_message", "prep_error", "prep_ready", "prep_has_mask_out"): |
| st.session_state.pop(k, None) |
| st.rerun() |
|
|
| t1_dir = MODEL_INPUT_DEST / "t1" |
| mask_dir = MODEL_INPUT_DEST / "masks" |
| prepared_t1 = sorted(t1_dir.glob("*.nii*")) if t1_dir.exists() else [] |
|
|
| if not prepared_t1: |
| st.info( |
| "No prepared files found yet. Run the registration button above to generate " |
| "`data/processed/test_input/t1` and optional masks." |
| ) |
| return |
|
|
| case_names = [p.name for p in prepared_t1] |
| selected_name = st.selectbox("Prepared case", case_names, key="prepared_case") |
| selected_t1 = t1_dir / selected_name |
| selected_mask = _paired_mask_path(selected_t1, mask_dir) |
|
|
| prepared_data, prepared_zooms = _load_nifti(str(selected_t1)) |
| prepared_mask_data = None |
| if selected_mask is not None and selected_mask.exists(): |
| prepared_mask_data, _ = _load_nifti(str(selected_mask)) |
|
|
| _render_overlay_from_arrays( |
| image_data=prepared_data, |
| image_zooms=prepared_zooms, |
| mask_data=prepared_mask_data, |
| key_prefix="prepared_view", |
| title="Prepared registered image", |
| ) |
|
|
| st.subheader("4) Side-by-side comparison") |
| if native_data is None: |
| st.info("Upload a native T1 above to enable side-by-side native vs registered comparison.") |
| else: |
| _render_side_by_side_comparison( |
| native_data=native_data, |
| native_mask=native_mask_data, |
| prepared_data=prepared_data, |
| prepared_mask=prepared_mask_data, |
| ) |
|
|
| st.subheader("5) Lesion segmentation") |
| render_inference_panel( |
| model_input_root=MODEL_INPUT_DEST, |
| project_root=PROJECT_ROOT, |
| model_dir_default=PROJECT_ROOT / "ARC_ATLAS_Model_V3", |
| output_root_default=MODEL_OUTPUT_DEST, |
| model_cache_root_default=HF_MODEL_CACHE_ROOT, |
| hf_model_repo_default=HF_MODEL_DEFAULT_REPO, |
| hf_model_revision_default=HF_MODEL_DEFAULT_REVISION, |
| hf_model_subdir_default=HF_MODEL_DEFAULT_SUBDIR, |
| ) |
|
|
| st.subheader("Prepared paths") |
| st.code( |
| "\n".join( |
| [ |
| f"T1 directory: {t1_dir}", |
| f"Mask directory: {mask_dir} (optional)", |
| f"Manifest: {MODEL_INPUT_DEST / 'manifest.csv'}", |
| f"Selected T1: {selected_t1}", |
| f"Selected mask: {selected_mask if selected_mask else 'None'}", |
| ] |
| ) |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| render_app() |
|
|