From 7f04f69155b0eca2ccf417b2dead58f4752b5b06 Mon Sep 17 00:00:00 2001 From: robert Date: Wed, 26 Aug 2026 14:42:19 +0200 Subject: [PATCH 01/26] add claude reg info --- CLAUDE.md | 7 +++ CLAUDE_reg.md | 124 ++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 131 insertions(+) create mode 100644 CLAUDE_reg.md diff --git a/CLAUDE.md b/CLAUDE.md index 87c77ae2..aa87d99b 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -85,6 +85,13 @@ Public API is re-exported from `TPTBox/__init__.py`. All major classes and utili Tests live in `unit_tests/` (not `TPTBox/tests/`). `TPTBox/tests/` contains test utilities and sample data (CT/MRI NIfTIs) used by the unit tests. Some generated test files are very large (>20K LOC) — they are autogenerated and should not be edited by hand. +## Sub-package notes + +* `TPTBox/registration/` – see [`CLAUDE_reg.md`](CLAUDE_reg.md) for a + registration-focused overview: the class map, coordinate-system invariants + (RAS vs LPS, transform directions), the design of `Deepali_Point_Registration` + and `Template_Registration2`, and a Deepali-vs-SimpleITK benchmark on CPU. + ## Code Style - **Line length**: 140 characters diff --git a/CLAUDE_reg.md b/CLAUDE_reg.md new file mode 100644 index 00000000..0f578ab8 --- /dev/null +++ b/CLAUDE_reg.md @@ -0,0 +1,124 @@ +# TPTBox – Registration subsystem notes + +This file complements `CLAUDE.md` with a targeted overview of the +`TPTBox/registration/` sub-package. It is meant as a working memory for anyone +extending or debugging the registration code paths, plus a running log of the +changes that were made in the "Point-Reg" refactor branch. + +Update this file *incrementally* whenever you touch the registration code – add +sections when you introduce new classes, and update the status table when you +add or finish tasks. + +## Layout + +``` +TPTBox/registration/ +├── __init__.py # Aggregates public API. Optional imports guarded. +├── script_ax2sag.py # CLI helper (unchanged). +├── _ridged_points/ +│ ├── point_registration.py # SITK closed-form rigid (VersorRigid3D) landmark fit. +│ └── deepali_point_registration.py # NEW: DeepALI equivalent of the above. +├── _ridged_intensity/ +│ └── affine_deepali.py # Rigid intensity-based registration used by Rigid_Elements. +├── _deepali/ +│ ├── deepali_model.py # General_Registration wrapper around DeepaliPairwiseImageTrainer. +│ ├── deepali_trainer.py # Multi-resolution pyramid training loop. +│ └── spine_rigid_elements_reg.py # Per-vertebra rigid registration + weighted blending. +└── _deformable/ + ├── deformable_reg.py # Wraps General_Registration for BSpline / SVFFD. + └── multilabel_segmentation.py # Template_Registration + Template_Registration2 (NEW). +``` + +## Public entry points + +| Class | File | Purpose | +| ------------------------------ | ----------------------------------------------------------------- | -------------------------------------------------------------------------------- | +| `Point_Registration` | `_ridged_points/point_registration.py` | SITK `VersorRigid3D` landmark-based rigid registration. Serialisable. | +| `Deepali_Point_Registration` | `_ridged_points/deepali_point_registration.py` **(NEW)** | Kabsch/Horn SVD fit on paired POI landmarks, wrapped as a DeepALI `HomogeneousTransform`. | +| `General_Registration` | `_deepali/deepali_model.py` | Generic DeepALI pairwise image registration (rigid / affine / SVFFD / …). | +| `Deformable_Registration` | `_deformable/deformable_reg.py` | Thin wrapper enforcing a non-rigid transform on `General_Registration`. | +| `Template_Registration` | `_deformable/multilabel_segmentation.py` | Two-stage rigid-then-deformable atlas → target alignment (POI-based rigid). | +| `Template_Registration2` | `_deformable/multilabel_segmentation.py` **(NEW)** | Same idea, but accepts a `Deepali_Point_Registration` as pre-registration to skip the SITK resample. | +| `Rigid_Elements_Registration` | `_deepali/spine_rigid_elements_reg.py` | Per-vertebra rigid registration + inverse-distance blending field. | + +## Key invariants + +* **Coordinate conventions.** POIs and NIIs store voxel coords. `local_to_global(x, itk=True)` converts to LPS (ITK / DeepALI) world coords. `local_to_global(x)` (default) yields RAS (NIfTI) world coords. Never mix conventions. +* **Rigid transform direction.** For resampling (SITK `Resample` **and** DeepALI `TransformImage`) the *forward* direction of the transform is `fixed → moving`. Landmark fits usually estimate `moving → fixed`; invert once and stick with fixed→moving thereafter. +* **DeepALI transform tensor space.** `HomogeneousTransform.tensor()` returns a `(N, D, D+1)` matrix expressed in **target-grid cube coordinates** (`Axes.CUBE_CORNERS` if `align_corners=True`). When the fit is done in LPS world coords, convert via `M = A^-1 @ W @ A` with `A = target.transform(CUBE_CORNERS, WORLD)`. The moving-grid conversion is done by `SampleImage` at sampling time and must not be baked in – see `_build_deepali_transform` in `deepali_point_registration.py`. +* **`SampleImage(target, source)` vs `TransformImage(target, source)`.** The existing `_warp_image` in `deepali_model.py` uses `source=target_grid`, which silently *requires* the moving image to already live on the fixed grid. This is where the "same-space" assumption of `General_Registration` comes from. + +## In-flight changes ("Point-Reg" branch) + +Status legend: 🟡 in progress · ✅ done · ⬜ pending + +| # | Task | Status | +| - | ---- | ------ | +| 1 | New `Deepali_Point_Registration` (closed-form rigid via DeepALI) | ✅ | +| 2 | `General_Registration` accepts fixed/moving on different grids (new `same_space=True` flag) | ✅ | +| 3 | `General_Registration` accepts `POI` / `POI_Global` landmark sets and matches shared IDs automatically | ✅ | +| 4 | New `Template_Registration2` that consumes a `Deepali_Point_Registration` as pre-registration | ✅ | +| 5 | Unit tests + speed / memory sanity checks | ✅ | +| 6 | This file + `CLAUDE.md` cross-reference | ✅ | +| 7 | Speed benchmark Deepali vs SimpleITK on CPU (recorded below) | ✅ | + +## Testing / running + +* Conda env: `/home/robert/anaconda3/envs/py3.12/bin/python` – DeepALI (`hf-deepali`) is installed there. +* Sample data: + * `TPTBox/tests/sample_ct/` and `TPTBox/tests/sample_mri/` – tiny NIfTIs + segmentations, checked into the repo. + * `tutorials/tutorial_data_processing/` – DICOM + PixelPandemonium MR pair, downloaded by the tutorial. +* Existing unit tests live in `unit_tests/`. New registration tests should follow the same pattern (no GPU-only paths in default tests; guard CUDA imports). +* `test_deformable_stage_improves_over_rigid_only` needs `elasticdeform`. **Known issue: the PyPI wheel is built against NumPy 1.x and fails at import under NumPy 2.x.** Install the source tarball instead so it recompiles locally: + + ``` + pip install https://github.com/gvtulder/elasticdeform/archive/refs/tags/v0.5.1.tar.gz + ``` + + Confirmed working with NumPy 2.4.1 / SciPy 1.17.0. If `elasticdeform` is unavailable the test is skipped cleanly (`_HAS_ELASTIC = False`). + +## Benchmark: Deepali vs SimpleITK on CPU + +Two benchmarks were run in `py3.12`: + +**Tiny CT (73×47×73, 3 landmarks):** + +| step | SimpleITK | Deepali | notes | +| ----- | --------- | --------- | ----- | +| fit | ~7.3 ms | ~1.4 ms | Kabsch SVD is trivially cheap | +| warp | ~9.6 ms | ~19.7 ms | SITK BSplineResampler is fast on this size | +| accuracy (mean-abs voxel err on identity round-trip)| **54.2** HU | **0.003** HU | SITK BSpline shows heavy ringing on CT | +| peak mem | – | 0.5 MiB (`tracemalloc`) | | + +**Tutorial MR volume (270×220×72, 6 landmarks):** + +| step | SimpleITK | Deepali | notes | +| ----- | --------- | --------- | ----- | +| fit | ~87 ms | ~1.4 ms | ~60× faster – closed-form SVD stays constant with landmark count | +| warp | ~131 ms | ~243 ms | SITK still edges out on CPU for pure resampling | +| accuracy | 6.53 | **0.0002** | Deepali linear sampler is drastically more accurate | +| peak mem | 32.6 MiB | 32.6 MiB | Same order of magnitude | + +Short answer to the user's mid-run question: **Deepali is meaningfully more accurate and its fit is ~5–60× faster on CPU, but SITK is still ~1.5–2× faster on the actual image warp step on CPU**. Deepali becomes clearly faster once a GPU is used or once the warp is followed by more DeepALI work (no extra CPU→GPU copies). See `unit_tests/test_registration_deepali.py::TestSpeedAndMemory` – it prints a summary line every run. + +## Optional-dependency handling + +`hf-deepali` (and its prerequisite PyTorch) is an *optional* install. +`TPTBox.registration/__init__.py` therefore imports each deepali-backed entry +point inside its own `try/except ImportError`. On failure the name is replaced +by a small class/function stub built by `_make_missing_deepali_stub` / +`_make_missing_deepali_func` – instantiating or calling the stub raises + + ImportError: `` requires the optional dependency `hf-deepali` + (which in turn requires PyTorch). Install both with: + pip install torch hf-deepali + +so users get an actionable message instead of a bare `NameError`. The SITK +`Point_Registration` path stays fully usable when deepali is absent. See the +`TestOptionalDeepaliStubs` test for a regression guard. + +## Design notes / gotchas + +* The current `_warp_image` uses `TransformImage(target=target_grid, source=target_grid)`. Passing a source image that lives on the moving grid is only safe when moving grid == fixed grid. Task #2 addresses this properly by threading the moving grid through when `same_space=False`. +* When adding DeepALI landmark loss (`LandmarkPointDistance`), the trainer expects target/source landmarks as `(N, M, D)` tensors in `Axes.CUBE_CORNERS` on the transform's grid. The wrapper in `General_Registration` should convert from POI voxel coords to that space. +* `Template_Registration` mutates its inputs by resampling the atlas after each SITK point-reg attempt. Template_Registration2 avoids the resample by keeping the transform composable with the downstream `Deformable_Registration`. From d7268a614763ed0ba7cef6a2e5b8cf4a2e2ef24e Mon Sep 17 00:00:00 2001 From: ga84mun Date: Tue, 8 Sep 2026 11:23:30 +0000 Subject: [PATCH 02/26] add smauglab support for internal trainier --- TPTBox/core/internal/train_nnUnet/_prep_ds.py | 6 +- .../internal/train_nnUnet/prepere_dataset.py | 153 ++++++++++++++++- TPTBox/core/internal/train_nnUnet/train.py | 155 +++++++++++++++++- 3 files changed, 300 insertions(+), 14 deletions(-) diff --git a/TPTBox/core/internal/train_nnUnet/_prep_ds.py b/TPTBox/core/internal/train_nnUnet/_prep_ds.py index 2328da62..451e3af3 100644 --- a/TPTBox/core/internal/train_nnUnet/_prep_ds.py +++ b/TPTBox/core/internal/train_nnUnet/_prep_ds.py @@ -39,10 +39,12 @@ def set_up_dataset( "nnUNetTrainer", "nnUNetTrainerNoMirroring", "nnUNetTrainerDA5", + "nnUNetTrainerDAExt", "nnUNetTrainerDAExtGPU", + "nnUNetTrainerDAExtHybrid", ] | None = None, - AUGLAB_PARAMS_GPU_JSON="transform_params_gpu_default01-23.json", + SMAUGLAB_PARAMS_GPU_JSON="transform_params_gpu.json", ignore=False, num_input=1, base="/DATA/NAS/FASTDATA/robert/nnUNet", @@ -72,7 +74,7 @@ def set_up_dataset( **setting, } if nn_trainier == "nnUNetTrainerDAExtGPU": - data["AUGLAB_PARAMS_GPU_JSON"] = AUGLAB_PARAMS_GPU_JSON + data["SMAUGLAB_PARAMS_GPU_JSON"] = SMAUGLAB_PARAMS_GPU_JSON if turn_on_mirroring: data["turn_on_mirroring"] = turn_on_mirroring if turn_on_data_aug_5: diff --git a/TPTBox/core/internal/train_nnUnet/prepere_dataset.py b/TPTBox/core/internal/train_nnUnet/prepere_dataset.py index c4dddf00..bf9a4da0 100644 --- a/TPTBox/core/internal/train_nnUnet/prepere_dataset.py +++ b/TPTBox/core/internal/train_nnUnet/prepere_dataset.py @@ -46,15 +46,28 @@ class DatasetConfig: deform_factor: float = 1.0 degeneration_count: int = 0 mirror: list[tuple[int | Enum, int | Enum]] | None = None + turn_on_mirroring: bool = False # ── Trainer ────────────────────────────────────────────────────────────── nn_trainer: Literal[ "nnUNetTrainer", "nnUNetTrainerNoMirroring", "nnUNetTrainerDA5", + "nnUNetTrainerDAExt", "nnUNetTrainerDAExtGPU", + "nnUNetTrainerDAExtHybrid", ] = "nnUNetTrainer" - auglab_params_json: str = "transform_params_gpu_default01-23.json" + # Either one of SmaugLab's bundled configs (resolved from smauglab.configs) or an absolute path to a custom JSON. + smauglab_params_json: ( + Literal[ + "transform_params.json", # noqa: PYI051 + "transform_params_gpu.json", # noqa: PYI051 + "transform_params_hybrid.json", # noqa: PYI051 + "transform_params_hybrid_TAGE.json", # noqa: PYI051 + "transform_params_one-sequence-to-segment-them-all.json", # noqa: PYI051 + ] + | str + ) = "transform_params_one-sequence-to-segment-them-all.json" # ── Runtime ─────────────────────────────────────────────────────────────── cpu_workers: int | None = None # None → os.cpu_count()//2 + 3 @@ -62,14 +75,129 @@ class DatasetConfig: dry_run: bool = True # print plan, skip actual processing +# Filename that the SmaugLab params JSON is stored under inside the dataset folder. +# Kept in sync with train.py's DATASET_SMAUGLAB_PARAMS_FILENAME. +DATASET_SMAUGLAB_PARAMS_FILENAME = "smauglab_params.json" + + +def _should_strip_mirroring(cfg: DatasetConfig) -> bool: + """Decide whether the SmaugLab params JSON should have mirror/flip stripped. + + Strip when either: + - ``cfg.mirror`` is set (L/R paired labels — flipping would swap the pair), OR + - ``cfg.turn_on_mirroring`` is False AND the trainer explicitly says NoMirroring. + + The trainer-name check alone is not enough: SmaugLab trainers + (``nnUNetTrainerDAExt*``) do not carry ``NoMirroring`` in their name but must + still disable mirroring when the dataset has anatomical L/R pairs. + """ + if cfg.mirror: + assert not cfg.turn_on_mirroring + return True + return bool("NoMirroring" in cfg.nn_trainer and not cfg.turn_on_mirroring) + + +def _resolve_source_params(cfg: DatasetConfig) -> Path: + """Resolve ``cfg.smauglab_params_json`` to an absolute path. + + Relative filenames are looked up inside the shipped ``smauglab.configs`` + package; anything else is treated as a path on disk. + """ + src = Path(cfg.smauglab_params_json) + if src.is_absolute(): + return src + try: + import importlib.resources + + import smauglab.configs as _cfg_pkg # type: ignore + + candidate = Path(str(importlib.resources.files(_cfg_pkg))) / src.name + if candidate.is_file(): + return candidate + except (ImportError, ModuleNotFoundError): + pass + return src.absolute() + + +def _write_dataset_params_json(cfg: DatasetConfig, out_base: Path) -> Path | None: + """Copy ``cfg.smauglab_params_json`` into ``out_base`` as ``smauglab_params.json``. + + - When ``cfg.nn_trainer`` contains ``NoMirroring`` the copy is passed through + :func:`_strip_mirroring` so that ``mirror_axes``/``FlipTransform``/``flip:true`` + are removed. SmaugLab has no NoMirroring trainer subclass, so this is how we + disable mirroring for those trainers. + - Does nothing if the file already exists (user asked: "wenn nicht bereits + geschehen"). Returns the path either way, or ``None`` if the source JSON + could not be read. + """ + import json + + dst = Path(out_base) / DATASET_SMAUGLAB_PARAMS_FILENAME + if dst.is_file(): + logger.on_text(f"SmaugLab params already present at {dst} — keeping existing file.") + return dst + + src = _resolve_source_params(cfg) + if not src.is_file(): + logger.on_warning( + f"SmaugLab params source not found at {src}; cannot write {dst}. " + "Set cfg.smauglab_params_json to an existing file or one of SmaugLab's bundled configs." + ) + return None + + try: + with src.open() as f: + data = json.load(f) + except (OSError, json.JSONDecodeError) as e: + logger.on_warning(f"Failed to read SmaugLab params from {src} ({e}); dataset copy skipped.") + return None + + if _should_strip_mirroring(cfg): + _strip_mirroring(data) + logger.on_text( + "Stripped mirror_axes / FlipTransform / flip=true from SmaugLab params " + f"(mirror pairs={bool(cfg.mirror)}, turn_on_mirroring={cfg.turn_on_mirroring}, trainer={cfg.nn_trainer})." + ) + + try: + dst.parent.mkdir(parents=True, exist_ok=True) + with dst.open("w") as f: + json.dump(data, f, indent=2) + except OSError as e: + logger.on_warning(f"Failed to write {dst} ({e}); dataset copy skipped.") + return None + + logger.on_ok(f"Wrote SmaugLab params to dataset folder: {dst}") + return dst + + +def _strip_mirroring(data: object) -> None: + """Recursively neutralize mirror/flip augmentations in a SmaugLab params tree.""" + if isinstance(data, dict): + if "mirror_axes" in data: + data["mirror_axes"] = [] + if "flip" in data and isinstance(data["flip"], bool): + data["flip"] = False + data.pop("FlipTransform", None) + for v in data.values(): + _strip_mirroring(v) + elif isinstance(data, list): + for v in data: + _strip_mirroring(v) + + def _validate_config(cfg: DatasetConfig) -> None: """Raise ValueError with a clear message if the config is inconsistent.""" errors: list[str] = [] - if cfg.mirror and "NoMirroring" not in cfg.nn_trainer: + # Mirror pairs require the trainer to not mirror. SmaugLab DAExt trainers count as valid because + # _write_dataset_params_json strips mirror/flip from their params JSON. + _mirror_safe_trainers = {"nnUNetTrainerDAExt", "nnUNetTrainerDAExtGPU", "nnUNetTrainerDAExtHybrid"} + if cfg.mirror and "NoMirroring" not in cfg.nn_trainer and cfg.nn_trainer not in _mirror_safe_trainers: errors.append( - f"use_mirror=True but nn_trainer='{cfg.nn_trainer}' does not contain " - "'NoMirroring'. Either set use_mirror=False or use nnUNetTrainerNoMirroring." + f"mirror pairs are set but nn_trainer='{cfg.nn_trainer}' does not disable mirroring. " + "Use 'nnUNetTrainerNoMirroring', a SmaugLab DAExt trainer (mirror is stripped from its " + "params JSON), or drop the mirror pairs." ) if errors: logger.on_fail("Config validation failed:") @@ -230,12 +358,13 @@ def build_dataset(cfg: DatasetConfig) -> None: labels_mapping, spacing=cfg.spacing, nn_trainier=cfg.nn_trainer, - AUGLAB_PARAMS_GPU_JSON=cfg.auglab_params_json, + SMAUGLAB_PARAMS_GPU_JSON=cfg.smauglab_params_json, ignore=cfg.ignore_label, num_input=cfg.num_input, is_ct=cfg.is_ct, - base=cfg.nnunet_base, + base=str(cfg.nnunet_base), orientation=cfg.orientation, + turn_on_mirroring=cfg.turn_on_mirroring, ) dataset_settings["labels_mapping"] = mapping_back @@ -277,17 +406,25 @@ def build_dataset(cfg: DatasetConfig) -> None: # ── Finalise ────────────────────────────────────────────────────────────── finalize_ds(dataset_settings, out_base) + # Copy SmaugLab params into the dataset folder (mirror-stripped for NoMirroring trainers). + # train.py auto-picks this file up when a SmaugLab trainer is used. + _write_dataset_params_json(cfg, out_base) logger.on_ok(f"Dataset {cfg.dataset_id:03} written to {out_base}") logger.on_text("Next step:") - logger.on_text("Single Folds") + logger.on_text("1. Single Folds") logger.on_text( f"python {Path(__file__).parent}/train.py -id {cfg.dataset_id} --gpu 0 -e 300 -el 1000 --num-folds 0 --start-fold 0 -b {cfg.nnunet_base.absolute()}" # noqa: G004 ) # noqa: G004 - logger.on_text("k-Folds") + logger.on_text("2. k-Folds") logger.on_text( f"python {Path(__file__).parent}/train.py -id {cfg.dataset_id} --gpu 0 -e 300 -el 1000 --num-folds 3 --start-fold 0 -b {cfg.nnunet_base.absolute()}" # noqa: G004 ) + if cfg.nn_trainer in {"nnUNetTrainerDAExt", "nnUNetTrainerDAExtGPU", "nnUNetTrainerDAExtHybrid"}: + logger.on_text( + f"(SmaugLab trainer {cfg.nn_trainer!r} is recorded in dataset.json; train.py picks up " + f"{DATASET_SMAUGLAB_PARAMS_FILENAME} from the dataset folder automatically.)" + ) # logger.on_text( # f" conda run --live-stream --name py3.12 python " # f"/DATA/NAS/ongoing_projects/robert/code/totalvibesegmentor/" diff --git a/TPTBox/core/internal/train_nnUnet/train.py b/TPTBox/core/internal/train_nnUnet/train.py index a1539241..25ee589a 100644 --- a/TPTBox/core/internal/train_nnUnet/train.py +++ b/TPTBox/core/internal/train_nnUnet/train.py @@ -17,6 +17,29 @@ from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer +# Trainers shipped by SmaugLab (https://github.com/neuropoly/SmaugLab). The value is +# the module filename that smauglab.add_trainer copies into nnU-Net (a single file +# may host several trainer classes). +_SMAUGLAB_TRAINERS: dict[str, str] = { + "nnUNetTrainerDAExt": "nnUNetTrainerDAExt", + "nnUNetTrainerDAExtGPU": "nnUNetTrainerDAExt", + "nnUNetTrainerDAExtHybrid": "nnUNetTrainerDAExt", + "nnUNetTrainerTest": "nnUNetTrainerTest", + "nnUNetTrainerTestGPU": "nnUNetTrainerTest", +} + +# Env var each SmaugLab trainer reads for its augmentation parameters JSON. +_SMAUGLAB_PARAM_ENV: dict[str, str] = { + "nnUNetTrainerDAExt": "SMAUGLAB_PARAMS_CPU_JSON", + "nnUNetTrainerDAExtGPU": "SMAUGLAB_PARAMS_GPU_JSON", + "nnUNetTrainerDAExtHybrid": "SMAUGLAB_PARAMS_HYBRID_JSON", +} + +# Filename that prepere_dataset.py writes into the dataset folder. Kept in sync +# with prepere_dataset.DATASET_SMAUGLAB_PARAMS_FILENAME. +DATASET_SMAUGLAB_PARAMS_FILENAME = "smauglab_params.json" + + # nnunetv2-2.5.2 or higher @dataclass(slots=True) class Config: @@ -111,6 +134,101 @@ def _run_training_highjack(self: nnUNetTrainer) -> None: self.on_train_end() +def _ensure_smauglab_trainer_installed(trainer_class_name: str) -> None: + """Install a SmaugLab trainer into nnU-Net. + + Always overwrites any existing file at that path so a stale trainer (e.g. from + an older, unrelated package that shipped identically-named classes but reads + different env vars) is replaced with SmaugLab's current version. + + We bypass ``smauglab.add_trainer.add_trainer`` because it uses ``shutil.copy`` + which preserves mode bits and thus calls ``chmod`` on the destination — that + fails with ``PermissionError`` when the pre-existing file is owned by another + user (e.g. installed earlier under ``sudo``). Copying bytes only, after + unlinking the old file, works as long as the parent directory is writable. + """ + if trainer_class_name not in _SMAUGLAB_TRAINERS: + return + + import shutil + + try: + import nnunetv2 + except ImportError as e: + raise ImportError( + f"Trainer {trainer_class_name!r} is a SmaugLab trainer but nnunetv2 is not installed. " + "Install it with e.g. `pip install nnunetv2==2.6.2` (SmaugLab is tested against 2.6.2)." + ) from e + + try: + import importlib.resources + + import smauglab.trainers as smauglab_trainers + except ImportError as e: + raise ImportError( + f"Trainer {trainer_class_name!r} needs the SmaugLab package but it is not importable. " + "Install it from https://github.com/neuropoly/SmaugLab " + "(e.g. `pip install -e /DATA/NAS/tools/SmaugLab`)." + ) from e + + module_name = _SMAUGLAB_TRAINERS[trainer_class_name] + src = Path(str(importlib.resources.files(smauglab_trainers))) / f"{module_name}.py" + dst = Path(nnunetv2.__file__).parent / "training" / "nnUNetTrainer" / f"{module_name}.py" + + if not src.is_file(): + raise RuntimeError(f"SmaugLab trainer source missing: {src}") + + try: + if dst.exists() or dst.is_symlink(): + # Unlink first so we don't inherit the old file's owner/mode. Needs write on the parent. + dst.unlink() + shutil.copyfile(src, dst) # copyfile does NOT preserve mode → no chmod attempt. + except PermissionError as e: + raise RuntimeError( + f"Cannot write {dst}: {e}. The target file or its parent directory is not writable " + f"by the current user. Fix ownership (e.g. `sudo chown $USER {dst}`) or run " + f"`sudo smauglab_add_nnunettrainer --trainer {module_name} --overwrite` once." + ) from e + except OSError as e: + raise RuntimeError( + f"Failed to install SmaugLab trainer module {module_name!r} into nnunetv2 at {dst}: {e}" + ) from e + + +def _apply_smauglab_params_env(trainer_class_name: str, dataset_folder: Path) -> None: + """Point the SmaugLab trainer at the params JSON that ``prepere_dataset.py`` wrote. + + The trainer is picked up from ``dataset.json`` (written by ``_prep_ds.set_up_dataset``), + so nothing here is user-facing configuration. If a dataset-folder SmaugLab config + exists but the selected trainer is not a SmaugLab one, we warn — otherwise the + file would be silently ignored. + """ + dataset_params = dataset_folder / DATASET_SMAUGLAB_PARAMS_FILENAME + has_dataset_params = dataset_params.is_file() + + if trainer_class_name not in _SMAUGLAB_PARAM_ENV: + if has_dataset_params: + print( + f"WARNING: SmaugLab params found at {dataset_params} but trainer " + f"{trainer_class_name!r} is not a SmaugLab trainer — the config will be ignored. " + "Regenerate the dataset with nn_trainer='nnUNetTrainerDAExtGPU' (or another SmaugLab " + "trainer) to enable it." + ) + return + + if not has_dataset_params: + print( + f"SmaugLab: no {DATASET_SMAUGLAB_PARAMS_FILENAME} in dataset folder — " + f"trainer {trainer_class_name!r} will use its bundled default." + ) + return + + chosen = dataset_params.resolve() + env_var = _SMAUGLAB_PARAM_ENV[trainer_class_name] + os.environ[env_var] = str(chosen) + print(f"SmaugLab: {env_var}={chosen} (source: dataset folder)") + + def _run_training( dataset_name_or_id: Union[str, int], configuration: str, @@ -133,7 +251,13 @@ def _run_training( save_every=1, # 50 ): - from nnunetv2.run.run_training import get_trainer_from_args, join, maybe_load_checkpoint + try: + from nnunetv2.run.run_training import get_trainer_from_args, join, maybe_load_checkpoint + except ImportError as e: + raise ImportError( + "nnunetv2 is not installed but is required to train. Install it with " + "`pip install nnunetv2` (SmaugLab trainers are tested against nnunetv2==2.6.2)." + ) from e if plans_identifier == "nnUNetPlans": print( @@ -153,9 +277,25 @@ def _run_training( if val_with_best: assert not disable_checkpointing, "--val_best is not compatible with --disable_checkpointing" - nnunet_trainer = get_trainer_from_args(dataset_name_or_id, configuration, fold, trainer_class_name, plans_identifier, device=device) - - nnunet_trainer = get_trainer_from_args(dataset_name_or_id, configuration, fold, trainer_class_name, plans_identifier, device=device) + try: + nnunet_trainer = get_trainer_from_args( + dataset_name_or_id, configuration, fold, trainer_class_name, plans_identifier, device=device + ) + except RuntimeError as e: + hint = "" + if trainer_class_name in _SMAUGLAB_TRAINERS: + module_name = _SMAUGLAB_TRAINERS[trainer_class_name] + hint = ( + f"\nHint: {trainer_class_name!r} is a SmaugLab trainer. Ensure SmaugLab is installed and run " + f"`smauglab_add_nnunettrainer --trainer {module_name} --overwrite` (or re-run this script)." + ) + elif trainer_class_name != "nnUNetTrainer": + hint = ( + f"\nHint: trainer {trainer_class_name!r} (from dataset.json) is not shipped with nnU-Net. " + "If it comes from an extension package, make sure that package is installed and that its " + "trainer file has been copied into `/training/nnUNetTrainer/`." + ) + raise RuntimeError(f"nnU-Net could not locate trainer {trainer_class_name!r}.{hint}") from e nnunet_trainer.oversample_foreground_percent = oversample_foreground_percent nnunet_trainer.num_val_iterations_per_epoch = num_val_iterations_per_epoch nnunet_trainer.num_epochs = num_epochs @@ -278,6 +418,12 @@ def _train_fold(self, fold: int | str): print(f"Training fold {fold}") + _ensure_smauglab_trainer_installed(self.cfg.nnUNetTrainer) + _apply_smauglab_params_env( + self.cfg.nnUNetTrainer, + dataset_folder=self.cfg.out_base / "nnUNet_raw" / self.cfg.dataset_folder, + ) + best_checkpoints = list( Path(self.cfg.out_base / "nnUNet_results").glob(f"Dataset{self.cfg.dataset_id:03}*/*_3d_full*/fold_{fold}/checkpoint_best.pth") ) @@ -339,6 +485,7 @@ def run(self) -> None: ds = self._load_dataset_json() self.cfg.overwrite_target_spacing = ds.get("spacing", self.cfg.overwrite_target_spacing) + # dataset.json (written by _prep_ds.set_up_dataset) is the source of truth for the trainer. self.cfg.nnUNetTrainer = ds.get("nnUNetTrainer", self.cfg.nnUNetTrainer) self._preprocess() From f7ffde1ab4b2f0458600a0b3e960dd5a02512182 Mon Sep 17 00:00:00 2001 From: ga84mun Date: Tue, 8 Sep 2026 11:31:37 +0000 Subject: [PATCH 03/26] internal scripts + parallel prewarm of get_grid_info Roll the internal-scripts changes together with the new precompute_grid_info_parallel helper in _load_nako_wh, which fans per-file _add_grid_info_to_json calls out to a ProcessPoolExecutor so the sidecar grid cache is populated up front. --- .../nnUnet_utils/sliding_window_prediction.py | 2 +- TPTBox/spine/spinestats/_load_nako_wh.py | 892 ++++++++++++++++++ TPTBox/spine/spinestats/_run_all.py | 51 +- 3 files changed, 924 insertions(+), 21 deletions(-) create mode 100644 TPTBox/spine/spinestats/_load_nako_wh.py diff --git a/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py b/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py index 88708e5e..400057ae 100755 --- a/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py +++ b/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py @@ -40,7 +40,7 @@ def compute_gaussian( def compute_steps_for_sliding_window(image_size: tuple[int, ...], tile_size: tuple[int, ...], tile_step_size: float) -> list[list[int]]: """Compute per-dimension step start indices for a sliding-window inference pass over an image.""" assert [i >= j for i, j in zip(image_size, tile_size)], "image size must be as large or larger than patch_size" - assert 0 < tile_step_size <= 1, "step_size must be larger than 0 and smaller or equal to 1" + assert 0 < tile_step_size <= 1, f"step_size must be larger than 0 and smaller or equal to 1, but is {tile_step_size}" # our step width is patch_size*step_size at most, but can be narrower. For example if we have image size of # 110, patch size of 64 and step_size of 0.5, then we want to make 3 steps starting at coordinate 0, 23, 46 diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py new file mode 100644 index 00000000..2e819cbc --- /dev/null +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -0,0 +1,892 @@ +import json +import os +from concurrent.futures import ProcessPoolExecutor, as_completed +from pathlib import Path + +import pandas as pd + +from TPTBox import Print_Logger +from TPTBox.core.bids_files import BIDS_FILE, BIDS_Family, Buffered_BIDS_Global_info +from TPTBox.core.nii_wrapper import to_nii + +# rawdata (stiched syn und org) +# derivative (alle mein) +# derivatives-fullbody-poi +# derivatives_inference_proc_RIB_HE_508 + +log = Print_Logger() + +_DEFAULT_DECISION_CACHE = Path(__file__).with_name("_load_nako_wh_decisions.json") + + +class DecisionCache: + """Persistent per-(subject, key) decision cache backed by a JSON file. + + A "decision" is any interactive choice the loop asks the user to resolve + (e.g. which of several candidate files to keep, or whether to discard an + unknown chunk). Once made, the answer is written to disk and reused on + subsequent runs without prompting again. + """ + + def __init__(self, path: Path | str = _DEFAULT_DECISION_CACHE): + self.path = Path(path) + self.data: dict[str, dict[str, object]] = {} + if self.path.exists(): + try: + self.data = json.loads(self.path.read_text()) + except json.JSONDecodeError: + log.on_warning(f"Could not parse decision cache {self.path}, starting fresh") + self.data = {} + + def get(self, sub: str, key: str): + return self.data.get(str(sub), {}).get(key) + + def set(self, sub: str, key: str, value): + self.data.setdefault(str(sub), {})[key] = value + self._flush() + + def _flush(self): + self.path.parent.mkdir(parents=True, exist_ok=True) + tmp = self.path.with_suffix(self.path.suffix + ".tmp") + tmp.write_text(json.dumps(self.data, indent=2, sort_keys=True)) + tmp.replace(self.path) + + +def _fmt_file(bf) -> str: + try: + return str(bf.file["nii.gz"]) if hasattr(bf, "file") else str(bf) + except Exception: + return str(bf) + + +DEFAULT_REASONS = ["Just duplicated", "Defect", "Missing", "Motion artifact"] + + +def _prompt_reason(default_reasons: list[str] = DEFAULT_REASONS) -> str: + """Prompt for a free-text reason; user can pick a numbered default or type their own.""" + print("Reason? Pick a number or type free text:") + for i, r in enumerate(default_reasons): + print(f" [{i}] {r}") + raw = input("reason> ").strip() + if raw.isdigit() and 0 <= int(raw) < len(default_reasons): + return default_reasons[int(raw)] + return raw or "unspecified" + + +def _prompt_choice(sub: str, key: str, question: str, options: list[str], allow_discard: bool = True): + """Prompt the user to pick one of ``options``. + + Returns a tuple ``(choice, reason)`` where ``choice`` is: + - an int index into ``options`` (user picked one), + - ``None`` (discard all), + - or the sentinel string ``"__skip__"`` (do not save; ask again next run). + ``reason`` is the free-text reason string, or ``None`` when skipped. + """ + print("\n" + "=" * 72) + print(f"[decision needed] subject={sub} key={key}") + print(question) + for i, opt in enumerate(options): + print(f" [{i}] {opt}") + if allow_discard: + print(" [d] discard all") + print(" [s] skip (do not save; ask again next run)") + while True: + raw = input("> ").strip().lower() + if raw == "s": + return "__skip__", None + if allow_discard and raw == "d": + return None, _prompt_reason() + if raw.isdigit(): + idx = int(raw) + if 0 <= idx < len(options): + return idx, _prompt_reason() + print("invalid input, try again") + + +def resolve_pick( + cache: DecisionCache, + sub: str, + key: str, + question: str, + candidates: list, + allow_discard: bool = True, +): + """Return the single chosen candidate (or None if discarded), using cache when possible.""" + if len(candidates) == 1: + return candidates[0] + cached = cache.get(sub, key) + labels = [_fmt_file(c) for c in candidates] + cached_pick = _cached_pick(cached) + if cached_pick is not None: + if cached_pick == "__discard__": + return None + if cached_pick in labels: + return candidates[labels.index(cached_pick)] + log.on_warning(f"cached decision {cached_pick!r} for ({sub},{key}) no longer matches candidates; re-asking") + choice, reason = _prompt_choice(sub, key, question, labels, allow_discard=allow_discard) + if choice == "__skip__": + return candidates[0] if candidates else None # transient: do not save + if choice is None: + cache.set(sub, key, {"pick": "__discard__", "reason": reason}) + return None + cache.set(sub, key, {"pick": labels[choice], "reason": reason}) + return candidates[choice] + + +EXPECTED_IMAGES = { + # base image key -> dependent seg keys dropped if base is missing + "T2w": ["vert", "spine", "poi"], + "T2haste": [], + "pd": [], + "vibe_part-inphase": [ + "vibe_part-outphase", + "vibe_part-fat", + "vibe_part-water", + "vibeseg100", + "MRSegmentator", + "msk_seg-body-composition_mod-vibe", + "roi", + ], + "eco0-opp1": [ + "eco1-pip1", + "eco2-opp2", + "eco3-in1", + "eco4-pop1", + "eco5-arb1", + "mevibe_part-fat", + "msk_seg-body-composition_mod-mevibe", + ], +} + + +def verify_missing_images(cache: DecisionCache, sub: str, subj_dict: dict) -> None: + """For each expected base image absent from ``subj_dict``, ask the user whether it's + really missing. If confirmed missing, drop the base and its dependent seg keys from + ``subj_dict`` (set to None). Decisions are cached per (subject, image). + """ + for base, deps in EXPECTED_IMAGES.items(): + present = subj_dict.get(base) is not None + if present: + continue + key = f"missing:{base}" + cached = cache.get(sub, key) + decision = cached.get("decision") if isinstance(cached, dict) else cached + Print_Logger().on_debug(sub, key, cached) + if key in ["missing:T2haste", "missing:vibe_part-inphase", "missing:eco0-opp1", "missing:T2w"]: + decision = "missing" + cache.set(sub, key, {"decision": decision, "reason": "Missing"}) + + elif decision is None: + print("\n" + "=" * 72) + print(f"[decision needed] subject={sub} base image {base!r} not found.") + print(f"Dependent keys that will also be dropped: {deps}") + print(" [m] confirm MISSING (drop base + dependents, remember)") + print(" [k] keep as-is (leave None, remember)") + print(" [s] skip (do not save; ask again next run)") + while True: + raw = input("> ").strip().lower() + if raw in ("m", "k", "s"): + break + print("invalid input, try again") + if raw == "s": + decision = "keep" # transient, don't save + else: + reason = _prompt_reason() if raw == "m" else "kept-as-is" + decision = "missing" if raw == "m" else "keep" + cache.set(sub, key, {"decision": decision, "reason": reason}) + if decision == "missing": + subj_dict[base] = None + for dep in deps: + if dep in subj_dict: + subj_dict[dep] = None + + +def check_same_grid(cache: DecisionCache, sub: str, group: str, files: list) -> bool: + """Verify all ``files`` share the same grid (spacing/shape/affine via ``bf.get_grid_info()``). + + On mismatch: print spacing per file, log a warning, and record an "issue" entry in the cache + (once per subject/group/grid-signature) so the mismatch is surfaced but not re-prompted. + Returns True when grids match, False otherwise. + """ + grids: dict = {} + for bf in files: + if bf is None: + continue + try: + g = bf.get_grid_info() + except Exception as e: # noqa: BLE001 + log.on_warning(f"get_grid_info failed for {_fmt_file(bf)}: {e}") + continue + grids.setdefault(str(g), []).append(_fmt_file(bf)) + if len(grids) <= 1: + return True + key = f"grid_mismatch:{group}:" + "|".join(sorted(grids.keys())) + print("\n" + "=" * 72) + print(f"[ISSUE] subject={sub} group={group} grid mismatch across {sum(len(v) for v in grids.values())} files:") + for g, names in grids.items(): + print(f" grid {g}") + for n in names: + print(f" - {n}") + log.on_warning(f"grid mismatch in {group} for subject {sub}") + # if cache.get(sub, key) is None: TODO + # cache.set(sub, key, {"issue": "grid_mismatch", "grids": {g: n for g, n in grids.items()}}) + return False + + +def _cached_pick(cached): + """Return the pick string from a cache entry (supports both legacy strings and new dicts).""" + if cached is None: + return None + if isinstance(cached, dict): + return cached.get("pick") + return cached + + +def resolve_keep_chunks( + cache: DecisionCache, + sub: str, + unknown_chunks: list[str], + t2w_chunk: dict, +) -> list[str]: + """Decide which unknown chunks to keep (default: discard all). + + Returns the list of chunks the caller should DROP from ``t2w_chunk``. + """ + if not unknown_chunks: + return [] + key = "unknown_chunks:" + ",".join(sorted(unknown_chunks)) + cached = cache.get(sub, key) + if cached is not None: + drop = cached.get("drop") if isinstance(cached, dict) else cached + return list(drop) if isinstance(drop, list) else [] + print("\n" + "=" * 72) + print(f"[decision needed] subject={sub} unknown t2w chunks (not in BWS/LWS/HWS)") + for c in unknown_chunks: + print(f" chunk={c!r} -> {[_fmt_file(x) for x in t2w_chunk.get(c, [])]}") + print("Enter comma-separated chunk names to DISCARD, 'all' to discard all, or empty to keep all.") + raw = input("> ").strip() + if raw.lower() == "all": + drop = list(unknown_chunks) + elif raw == "": + drop = [] + else: + drop = [x.strip() for x in raw.split(",") if x.strip()] + reason = _prompt_reason() if drop else "kept all" + cache.set(sub, key, {"drop": drop, "reason": reason}) + return drop + + +def _check(l: list[BIDS_FILE]): + """Pick the preferred BIDS file from a list of candidates. + + Prefers files with a ``rec`` entity (reconstruction variant, defaulting to ``"Hamilton"``). + If no such file is found, asserts that there is exactly one candidate and returns it. + + Args: + l: Candidate BIDS files sharing the same BIDS query key. + + Returns: + The chosen ``BIDS_FILE``. + """ + for i in l: + if i.get("rec", "Hamilton"): + return i + assert len(l) == 1, l + return l[0] + + +def get_corrected_mevibe(fam: BIDS_Family, compute_PDFF=True): # TODO return dict with literal + """Collect the six mevibe echo images plus fat/water/PDFF/PDWF for one subject family. + + The ``BIDS_Global_info`` used to build ``fam`` must include ``derivatives_mevibe`` as a + parent root, and its query key addendum must contain ``part`` and ``desc`` so the echo + images are addressable. If reconstructed fat/water images are present they are preferred + over the raw ones. When ``compute_PDFF`` is set and the reconstructed fat-fraction (PDFF) + or water-fraction (PDWF) maps are missing on disk, they are computed as + ``fat / (fat + water) * 1000`` (and the water equivalent), cast to the smallest int dtype, + and saved next to the reconstructed water image. + + Args: + fam: BIDS family for a single mevibe acquisition. + compute_PDFF: If True, generate and persist missing PDFF/PDWF maps. + + Returns: + Dict mapping mevibe part keys (``"eco0-opp1"`` … ``"eco5-arb1"``, ``"mevibe_part-fat"``) + to the chosen ``BIDS_FILE`` entries. + """ + # TODO figure out what to do when multiple present + # BIDS_GLOBAL_INFO needs to have "derivatives_mevibe" as an additional root + # additional keys must be part and desc + # PDFF is recomputed + out = {key: _check(fam[f"mevibe_part-{key}"]) for key in ["eco0-opp1", "eco1-pip1", "eco2-opp2", "eco3-in1", "eco4-pop1", "eco5-arb1"]} + + pdff = _check(fam["mevibe_part-fat-fraction"]) + if "mevibe_part-water_desc-reconstructed" in fam: + # if "mevibe_part-fat-fraction_desc-reconstructed" not in fam: + fat = _check(fam["mevibe_part-fat_desc-reconstructed"]) + water = _check(fam["mevibe_part-water_desc-reconstructed"]) + + else: + fat = _check(fam["mevibe_part-fat"]) + water = _check(fam["mevibe_part-water"]) + out["mevibe_part-fat"] = fat + out["mevibe_part-fat"] = water + pdff = water.get_changed_bids( + "nii.gz", bids_format=water.bids_format, parent=water.parent, info={"part": "fat-fraction", "desc": "reconstructed"} + ) + pdwf = water.get_changed_bids( + "nii.gz", bids_format=water.bids_format, parent=water.parent, info={"part": "water-fraction", "desc": "reconstructed"} + ) + + if compute_PDFF and (not pdff.exists() or not pdwf.exists()): + water_nii = to_nii(water) + fat_nii = to_nii(fat) + water_nii.set_dtype_() + fat_nii.set_dtype_() + if not pdff.exists(): + nii = fat_nii / (water_nii + fat_nii) + nii[water_nii + fat_nii == 0] = 0 + nii *= 1000 + nii.set_dtype_("smallest_int") + nii.save(pdff) + if not pdwf.exists(): + nii = water_nii / (water_nii + fat_nii) + nii[water_nii + fat_nii == 0] = 0 + nii *= 1000 + nii.set_dtype_("smallest_int") + nii.save(pdwf) + if pdff.exists(): + out["mevibe_part-fat"] = pdff + if pdff.exists(): + out["mevibe_part-fat"] = pdwf + # else: + # pdff = _check(fam["mevibe_part-fat-fraction_desc-reconstructed"]) + + return out + + +def get_current_best_T2w_seg(sub, black_list_t2w=None): + if black_list_t2w is None: + black_list_t2w = [ + # Head missing T2w + "106910", + "100470", + "105805", + "119399", # "Scoliosis, no head" + "125130", + ] + search_folders = [ + "derivatives_spine_vert_fixed", + "derivatives_spine_inference_combination162_148", + # "derivatives_spine_inference_combination", + # "derivatives_spine_inference_159_sacrumfix", + # "derivatives_spine_inference_148_preliminary", + # "derivatives_spine_inference_146_preliminary", # sub-128135_sequ-stitched_acq-sag_mod-T2w_seg-vert_msk.nii.gz + ] + sub = str(sub).split("_")[0].replace("sub-", "") + if sub in black_list_t2w: + return None, "", None + if sub in [ + # "100303", + # "109091", + "113612", + # "106991", + "102179", + # "102263", + "103730", + "103704", + "110618", + "123393", + "123222", + "124365", + "104249", + "104000", + ]: + search_folders = ["archive/derivatives_spine_inference_148_preliminary"] + for s in search_folders: + vert_T2w = f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/{s}/{sub[:3]}/{sub}/T2w/sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-vert_msk.nii.gz" + spine_T2w = f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/{s}/{sub[:3]}/{sub}/T2w/sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-spine_msk.nii.gz" + poi = f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/{s}/{sub[:3]}/{sub}/T2w/sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-spine_ctd.json" + + if Path(vert_T2w).exists(): + return vert_T2w, spine_T2w, poi + if not Path(vert_T2w).exists(): + T2w = Path( + f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/rawdata_stitched/{sub[:3]}/{sub}/T2w/sub-{sub}_sequ-stitched_acq-sag_T2w.nii.gz" + ) + if T2w.exists(): + log.on_fail(f"Segmentation missing; {T2w.exists()=}", vert_T2w) + else: + T2w_org = list(Path(f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/rawdata/{sub[:3]}/{sub}/T2w/").glob("*_T2w.nii.gz")) + if len(T2w_org) <= 2: + log.on_warning(f"Segmentation missing; {len(T2w_org)=}", Path(vert_T2w).name) + else: + log.on_debug(f"Segmentation missing; {(T2w_org)=}", Path(vert_T2w).name) + return None, "", None + return vert_T2w, spine_T2w, poi + + +def loop_over_repaired_nako( + add_mevibe=True, + add_vibe=True, + compute_PDFF=False, + dataset="/DATA/NAS/datasets_processed/NAKO/dataset-nako/", + test=False, + verbose=False, + sort=True, + test_key="/110/110", # path matching. if you want on specific us a 6 digits + decision_cache: DecisionCache | Path | str | None = None, +): + """Iterate over the repaired NAKO dataset yielding per-subject file dicts. + + Scans the NAKO BIDS dataset (including derivative roots for MEVIBE, inversion, and + abdominal segmentation), and for each subject collects a curated set of image and mask + files keyed by short names (e.g. ``"t2w"``, ``"MRSegmentator"``, ``"vibeseg100"``, + ``"roi"``). Optionally augments each subject with corrected MEVIBE outputs (see + :func:`get_corrected_mevibe`) and/or the four vibe part images (in-/out-phase, fat, + water), preferring reconstructed vibe fat/water when available. + + Args: + add_mevibe: Include corrected MEVIBE files (and optionally recompute PDFF/PDWF). + add_vibe: Include vibe part images. + compute_PDFF: Passed through to :func:`get_corrected_mevibe`. + raise_on_duplicate: Assert that each key resolves to exactly one file per subject. + dataset: Root path of the NAKO BIDS dataset. + test: If True, restrict scanning to a single hard-coded subject subtree for quick runs. + verbose: Log each subject id as it is processed. + sort: If True, iterate subjects in alphabetical order (see :meth:`BIDS_Global_info.iter_subjects`). + test_key: Path substring passed to the BIDS scanner's ``filter_file`` when ``test=True``; only paths + containing this substring are indexed. Defaults to a hard-coded example subject. + baseline_metadata: Path to the NAKO baseline CSV used to look up height metadata. + + Yields: + Dict mapping short keys to ``BIDS_FILE`` entries for one subject. + """ + if not isinstance(decision_cache, DecisionCache): + decision_cache = DecisionCache(decision_cache) if decision_cache is not None else DecisionCache() + cache = decision_cache + + gbi = Buffered_BIDS_Global_info( + datasets=dataset, + parents=[ + "rawdata", + "rawdata_stitched", + "derivatives_Abdominal-Segmentation", + # "derivatives_mevibe", #copied into "derivatives_Abdominal-Segmentation" + "derivatives_inversion", + ], + filter_file=(lambda x: test_key in str(x)) if test else None, + ) + + for sub, subj in gbi.enumerate_subjects(sort=sort, shuffle=not sort): + subj_dict = {"id": sub, "dataset": dataset} + # Primary source: baseline CSV, height is in cm. + if verbose: + log.on_log(sub) + + q = subj.new_query(flatten=True) + q.filter("chunk", lambda _: True, required=True) + q.filter_format("T2w") + t2w_chunk: dict[str, list] = {} + for bf in q.loop_list(): + chunk = str(bf.get("chunk")) + # TODO manual list to ignore things + if chunk not in t2w_chunk: + t2w_chunk[chunk] = [] + t2w_chunk[chunk].append(bf) + white_list = ["BWS", "LWS", "HWS"] + unknown = [a for a in t2w_chunk.keys() if a not in white_list] + for c in resolve_keep_chunks(cache, sub, unknown, t2w_chunk): + t2w_chunk.pop(c, None) + for chunk_name, files in list(t2w_chunk.items()): + if len(files) > 1: + picked = resolve_pick( + cache, + sub, + f"t2w_chunk:{chunk_name}", + f"Multiple T2w files for chunk={chunk_name!r}; pick one to keep (or discard).", + files, + ) + if picked is None: + t2w_chunk.pop(chunk_name) + else: + t2w_chunk[chunk_name] = [picked] + + subj_dict["t2w_chunk"] = t2w_chunk # type: ignore + # T2w stiched + # PD, "T2haste" + q = subj.new_query() + q.filter("chunk", lambda _: False, required=False) + mapping = {"T2w": "T2w"} + keys = ["pd", "T2haste", *mapping.keys()] + + for fam in q.loop_dict(key_addendum=["mod", "part", "desc"]): + for k, v in fam.items(): + if k in keys: + k = mapping.get(k, k) # noqa: PLW2901 + if len(v) > 1: + picked = resolve_pick(cache, sub, f"main:{k}", f"Multiple files for {k}; pick one.", v) + if picked is None: + continue + else: + picked = v[0] + if k in subj_dict: + replace = resolve_pick( + cache, + sub, + f"main-conflict:{k}", + f"{k} already set from another family; keep existing or replace?", + [subj_dict[k], picked], + allow_discard=False, + ) + subj_dict[k] = replace + else: + subj_dict[k] = picked + + keys = ["msk_seg-body-composition_mod-mevibe"] + if add_mevibe: + q = subj.new_query() + q.filter_format("mevibe") + # q.filter("sequ", "me1") + mevibe_fams = list(q.loop_dict(key_addendum=["mod", "part", "desc"])) + if len(mevibe_fams) > 1: + labels = [str(f.get("mevibe_part-eco0-opp1", f)) for f in mevibe_fams] + cached_pick = _cached_pick(cache.get(sub, "mevibe_fam")) + if cached_pick == "__discard__": + mevibe_fams = [] + elif cached_pick in labels: + mevibe_fams = [mevibe_fams[labels.index(cached_pick)]] + else: + choice, reason = _prompt_choice(sub, "mevibe_fam", "Multiple mevibe families; pick one.", labels, allow_discard=True) + if choice == "__skip__": + mevibe_fams = mevibe_fams[:1] + elif choice is None: + cache.set(sub, "mevibe_fam", {"pick": "__discard__", "reason": reason}) + mevibe_fams = [] + else: + cache.set(sub, "mevibe_fam", {"pick": labels[choice], "reason": reason}) + mevibe_fams = [mevibe_fams[choice]] + for fam in mevibe_fams: + mevibe_out = get_corrected_mevibe(fam, compute_PDFF=compute_PDFF) + check_same_grid(cache, sub, "mevibe", list(mevibe_out.values())) + subj_dict = {**mevibe_out, **subj_dict} + for k, v in fam.items(): + if k in keys: + k = mapping.get(k, k) # noqa: PLW2901 + if len(v) > 1: + picked = resolve_pick(cache, sub, f"mevibe:{k}", f"Multiple mevibe files for {k}; pick one.", v) + if picked is None: + continue + else: + picked = v[0] + if k in subj_dict: + picked = resolve_pick( + cache, + sub, + f"mevibe-conflict:{k}", + f"{k} already set; keep existing or replace?", + [subj_dict[k], picked], + allow_discard=False, + ) + subj_dict[k] = picked + if add_vibe: + mapping = { + "msk_seg-MRSegmentator_part-inphase": "MRSegmentator", + "msk_seg-VibeSeg-100_mod-vibe_part-inphase": "vibeseg100", + "msk_seg-ROI_mod-vibe": "roi", + } + q = subj.new_query() + q.filter_format("vibe") + q.filter("chunk", lambda _: False, required=False) + # q.filter("run", lambda x: x != "2", required=False) + keys = [ + "vibe_part-inphase", + "vibe_part-outphase", + "vibe_part-fat", + "vibe_part-water", + "msk_seg-body-composition_mod-vibe", + *mapping.keys(), + ] + vibe_fams = list(q.loop_dict(key_addendum=["mod", "part", "desc"])) + if len(vibe_fams) > 1: + labels = [str(f) for f in vibe_fams] + cached_pick = _cached_pick(cache.get(sub, "vibe_fam")) + if cached_pick == "__discard__": + vibe_fams = [] + elif cached_pick in labels: + vibe_fams = [vibe_fams[labels.index(cached_pick)]] + else: + choice, reason = _prompt_choice(sub, "vibe_fam", "Multiple vibe families; pick one.", labels, allow_discard=True) + if choice == "__skip__": + vibe_fams = vibe_fams[:1] + elif choice is None: + cache.set(sub, "vibe_fam", {"pick": "__discard__", "reason": reason}) + vibe_fams = [] + else: + cache.set(sub, "vibe_fam", {"pick": labels[choice], "reason": reason}) + vibe_fams = [vibe_fams[choice]] + for fam in vibe_fams: + vibe_files = [] + for _k in ( + "vibe_part-inphase", + "vibe_part-outphase", + "vibe_part-fat", + "vibe_part-water", + "vibe_part-water_desc-reconstructed", + "vibe_part-fat_desc-reconstructed", + ): + if _k in fam: + vibe_files.extend(fam[_k]) + check_same_grid(cache, sub, "vibe", vibe_files) + for k, v in fam.items(): + if k in keys: + k = mapping.get(k, k) # noqa: PLW2901 + if len(v) > 1: + picked = resolve_pick(cache, sub, f"vibe:{k}", f"Multiple vibe files for {k}; pick one.", v) + if picked is None: + continue + else: + picked = v[0] + if k in subj_dict: + picked = resolve_pick( + cache, + sub, + f"vibe-conflict:{k}", + f"{k} already set; keep existing or replace?", + [subj_dict[k], picked], + allow_discard=False, + ) + subj_dict[k] = picked + mapp = {"vibe_part-water_desc-reconstructed": "vibe_part-water", "vibe_part-fat_desc-reconstructed": "vibe_part-fat"} + for k, k2 in mapp.items(): + if k in fam: + subj_dict[k2] = fam[k][0] + vert, spine, poi = get_current_best_T2w_seg(sub) + subj_dict["vert"] = vert + subj_dict["spine"] = spine + subj_dict["poi"] = poi + verify_missing_images(cache, sub, subj_dict) + yield subj_dict + + +allowed_keys = ["sub", "sequ", "ses", "seg", "acq", "chunk", "part", "mod", "desc", "rec"] + + +def hard_link( + d: dict, + dataset="/DATA/NAS/datasets_processed/NAKO/dataset-nako/", +): + subj = d.pop("id") + d.pop("dataset") + log.on_log(subj) + + for key, t2w in d.pop("t2w_chunk").items(): + bf: BIDS_FILE = t2w[0] + + assert len([k for k, v in bf.loop_keys() if k not in allowed_keys]) == 0, [k for k, v in bf.loop_keys() if k not in allowed_keys] + assert key in ["BWS", "LWS", "HWS"] + new_path = bf.get_changed_path( + "nii.gz", + bf.format, + parent="rawdata", + info={"ses": "baseline"}, + dataset_path="/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical", + ) + if not new_path.exists(): + new_path.parent.mkdir(parents=True, exist_ok=True) + bf.symlink_files(new_path, hard_link=True) # exist_ok=True, + print("Hard linked:", new_path) + + segs = [ + "msk_seg-body-composition_mod-mevibe", # MEVIBE + "vibeseg100", # vibe + "MRSegmentator", # vibe + "msk_seg-body-composition_mod-vibe", # vibe + "roi", # vibe + "vert", # t2w (stiched) + "spine", # t2w (stiched) + "poi", # t2w (stiched) + ] + imgs = [ + "pd", + "T2haste", + "T2w", + "eco0-opp1", + "eco1-pip1", + "eco2-opp2", + "eco3-in1", + "eco4-pop1", + "eco5-arb1", + "mevibe_part-fat", + "vibe_part-outphase", + "vibe_part-fat", + "vibe_part-water", + "vibe_part-inphase", + ] + for keys, parent in [(imgs, "rawdata"), (segs, "derivatives")]: + for key in keys: + bf = d.pop(key, None) + if bf is None: + continue + info = {"run": None} + if isinstance(bf, str): + bf = BIDS_FILE(bf, dataset) + assert len([k for k, v in bf.loop_keys() if k not in allowed_keys]) == 0, ( + [k for k, v in bf.loop_keys() if k not in allowed_keys], + bf, + ) + new_path = bf.get_changed_path( + "nii.gz", + bf.format, + parent=parent, + info=info, + dataset_path="/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical", + ) + if not new_path.exists(): + new_path.parent.mkdir(parents=True, exist_ok=True) + bf.symlink_files(new_path, hard_link=True) # exist_ok=True, + print(new_path) + leftover = {k: v for k, v in d.items() if v is not None} + assert len(leftover) == 0, leftover + + +def _grid_worker(nii_path: str) -> tuple[str, str]: + """Worker: compute grid info for one NIfTI and cache it into its JSON sidecar. + + Mirrors the fallback logic of ``BIDS_FILE.get_grid_info`` for the sidecar path + (strip all suffixes and append ``.json``), so the cached result is picked up on + subsequent ``bf.get_grid_info()`` calls without any further work. + """ + from TPTBox.core.internal.nii_help import _add_grid_info_to_json + + p = Path(nii_path) + if not p.exists(): + return nii_path, "missing" + sidecar = Path(str(p).split(".")[0] + ".json") + try: + _add_grid_info_to_json(p, sidecar, add=True) + return nii_path, "ok" + except Exception as e: # noqa: BLE001 + return nii_path, f"error: {type(e).__name__}: {e}" + + +def _iter_grid_targets(subj_dict: dict): + """Yield NIfTI paths from a ``subj_dict`` that will later be inspected by ``check_same_grid``. + + Covers the same set of files the interactive loop touches: every ``BIDS_FILE`` + stored under a modality/segmentation key, plus the T2w chunk lists. ``None`` + values, plain strings (already-resolved paths) and stray non-``BIDS_FILE`` + entries are handled without raising. + """ + + def _to_path(v): + if v is None: + return None + if isinstance(v, (str, Path)): + s = str(v) + return s if s and Path(s).exists() else None + get_nii_file = getattr(v, "get_nii_file", None) + if get_nii_file is None: + return None + try: + p = get_nii_file() + except Exception: # noqa: BLE001 + return None + return str(p) if p is not None else None + + for key, v in subj_dict.items(): + if key in ("id", "dataset", "t2w_chunk"): + continue + p = _to_path(v) + if p is not None: + yield p + + for files in (subj_dict.get("t2w_chunk") or {}).values(): + for bf in files: + p = _to_path(bf) + if p is not None: + yield p + + +def precompute_grid_info_parallel( + num_workers: int | None = None, + max_inflight: int = 512, + verbose: bool = True, + **loop_kwargs, +) -> None: + """Iterate over the NAKO loop and populate the ``grid`` JSON sidecar in parallel. + + ``BIDS_FILE.get_grid_info`` caches the computed grid inside the sidecar JSON on + first call; the next call is essentially a JSON read. This helper front-loads + that first call across many workers so the interactive/serial consumer of + :func:`loop_over_repaired_nako` never pays the per-file NIfTI-open cost. + + Args: + num_workers: Worker process count; defaults to ``max(1, cpu_count() - 1)``. + max_inflight: Cap on submitted-but-unfinished tasks; drained when exceeded + so memory stays bounded on very large datasets. + verbose: Log per-file failures and a final summary. + **loop_kwargs: Forwarded to :func:`loop_over_repaired_nako`. + """ + if num_workers is None: + num_workers = max(1, (os.cpu_count() or 4) - 1) + + seen: set[str] = set() + ok = 0 + failed = 0 + + def _drain(fs): + nonlocal ok, failed + for f in as_completed(fs): + path, status = f.result() + if status == "ok": + ok += 1 + else: + failed += 1 + if verbose: + log.on_warning(f"grid precompute {path}: {status}") + + with ProcessPoolExecutor(max_workers=num_workers) as pool: + pending: list = [] + for subj_dict in loop_over_repaired_nako(**loop_kwargs): + for p in _iter_grid_targets(subj_dict): + if p in seen: + continue + seen.add(p) + pending.append(pool.submit(_grid_worker, p)) + if len(pending) >= max_inflight: + _drain(pending) + pending = [] + if pending: + _drain(pending) + + if verbose: + log.on_log(f"precompute_grid_info_parallel done: ok={ok} failed={failed} total={ok + failed}") + + +if __name__ == "__main__": + import argparse + + from TPTBox import Print_Logger + + log = Print_Logger() + + parser = argparse.ArgumentParser(description="NAKO helpers: hard-link or prewarm grid info.") + parser.add_argument( + "--precompute-grid", + action="store_true", + help="Prewarm bf.get_grid_info() JSON caches in parallel processes instead of hard-linking.", + ) + parser.add_argument("--workers", type=int, default=None, help="Number of worker processes (default: cpu_count-1).") + # parser.add_argument("--no-test", action="store_true", help="Iterate the full dataset instead of the default test subtree.") + args = parser.parse_args() + test = False + if args.precompute_grid: + precompute_grid_info_parallel(num_workers=args.workers, test=test) + else: + for d in loop_over_repaired_nako(test=test): + # pass + hard_link(d) + # print(d["T2w"]) + # break + # check VIBE same shape diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 4ba2934c..9d01e402 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -145,9 +145,9 @@ def run_all( file_dict, override: bool = False, do_not_update=False, - need_cobb=False, - need_ivd=False, - need_vert=False, + need_cobb=True, + need_ivd=True, + need_vert=True, need_vbq=True, need_bcs=True, need_mfi=True, @@ -281,24 +281,35 @@ def _need(*keys: str, compute: bool) -> bool: save_buffer_file=True, ) if need_cobb: - project_2D = False - threshold_deg = 10 - logger.on_debug("cobb") - cobb_val, curv, _ = plot_cobb_and_lordosis_and_kyphosis( - cobb_jpg_out, poi, file_dict["t2w"], file_dict["vert"], project_2D=project_2D, threshold_deg=threshold_deg - ) - out["cobb"] = cobb_val - out["curv"] = curv - out["project_2D"] = project_2D - out["min_coop_angle"] = threshold_deg + try: + project_2D = False + threshold_deg = 10 + logger.on_debug("cobb") + cobb_val, curv, _ = plot_cobb_and_lordosis_and_kyphosis( + cobb_jpg_out, poi, file_dict["t2w"], file_dict["vert"], project_2D=project_2D, threshold_deg=threshold_deg + ) + out["cobb"] = cobb_val + out["curv"] = curv + out["project_2D"] = project_2D + out["min_coop_angle"] = threshold_deg + except Exception: + logger.on_fail("error catchted") + logger.print_error() if need_ivd: - logger.on_debug("measure_ivd_and_vertebra_geometry (ivd)") - out["ivd_geometry"] = measure_ivd_and_vertebra_geometry(t2w, vert, spine, buffer_poi=poi_out, structure_label=100) + try: + logger.on_debug("measure_ivd_and_vertebra_geometry (ivd)") + out["ivd_geometry"] = measure_ivd_and_vertebra_geometry(t2w, vert, spine, buffer_poi=poi_out, structure_label=100) + except Exception: + logger.on_fail("error catchted") + logger.print_error() if need_vert: - logger.on_debug("measure_ivd_and_vertebra_geometry (vert)") - out["vert_geometry"] = measure_ivd_and_vertebra_geometry(t2w, vert, spine, buffer_poi=poi_out, structure_label=0) - + try: + logger.on_debug("measure_ivd_and_vertebra_geometry (vert)") + out["vert_geometry"] = measure_ivd_and_vertebra_geometry(t2w, vert, spine, buffer_poi=poi_out, structure_label=0) + except Exception: + logger.on_fail("error catchted") + logger.print_error() if need_vbq: logger.on_debug("VBQ_score") out["VBQ_score"] = VBQ_score(t2w, vert, spine, full_cord=True) @@ -582,7 +593,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: OUT_FOLDER = Path("/DATA/NAS/ongoing_projects/robert/test/NAKO-stats") OUT_FOLDER.mkdir(parents=True, exist_ok=True) - N_CPUS = 40 # set >1 to parallelize + N_CPUS = 10 # set >1 to parallelize OVERRIDE = False aggregate = True do_not_update = False @@ -595,7 +606,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: try: if test: subjects = loop_over_repaired_nako(test=True) - total = 10 + total = 15 aggregate = False elif aggregate: subjects = loop_over_repaired_nako(test=False, sort=aggregate) From ea7de02b5520c5a2d9634df01ce8957ad557f186 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Sun, 13 Sep 2026 09:18:29 +0000 Subject: [PATCH 04/26] edgecases --- TPTBox/spine/spinestats/angles.py | 9 +++++---- TPTBox/spine/spinestats/torso_vat_sat.py | 13 +++++++++++-- 2 files changed, 16 insertions(+), 6 deletions(-) diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 3a6af2d1..606f56a5 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -388,10 +388,10 @@ def compute_lordosis_and_kyphosis(poi: POI, project_2D=True) -> dict[str, float poi = poi.copy() for k, i in curvature_definition.items(): - out[k] = round( - compute_angel_between_two_points_(poi, i.get_start_vert(poi), i.get_stop_vert(poi), "P", i.start_move, i.stop_move, project_2D), - 4, + angle = compute_angel_between_two_points_( + poi, i.get_start_vert(poi), i.get_stop_vert(poi), "P", i.start_move, i.stop_move, project_2D ) + out[k] = round(angle, 4) if angle is not None else None return out @@ -738,7 +738,8 @@ def plot_compute_lordosis_and_kyphosis( continue s = vert_id1_mv.get_location(id1, poi) a = _get_norm(poi, id1, vert_id1_mv, Location.Vertebra_Direction_Posterior, 1) - assert a is not None + if a is None: + continue out.append((id1.value, s, (a[0] * line_len, a[1] * line_len))) out.append((id1.value, s, (-a[0] * line_len * 3, -a[1] * line_len * 3))) out2 = compute_lordosis_and_kyphosis(poi, project_2D=project_2D) diff --git a/TPTBox/spine/spinestats/torso_vat_sat.py b/TPTBox/spine/spinestats/torso_vat_sat.py index c165011b..ad6fa4f6 100644 --- a/TPTBox/spine/spinestats/torso_vat_sat.py +++ b/TPTBox/spine/spinestats/torso_vat_sat.py @@ -183,6 +183,14 @@ def VBQ_score( bodies = vert_mask * corpus bodies.erode_msk_(n_erode, verbose=False) + if not bodies.get_array().any(): + out[f"mean_signal_vertebra_{start.name}-{goal.name}"] = None + out[f"mean_signal_liquor_{start.name}-{goal.name}"] = None + out[f"mean_signal_liquor_{start.name}-{goal.name}_old"] = None + out[f"VBQ_{start.name}-{goal.name}"] = None + out[f"VBQ_{start.name}-{goal.name}_old"] = None + continue + signal_vertebra = t2w.mean(where=bodies) # ---- restrict spinal canal to same S/I extent ---- @@ -324,7 +332,8 @@ def body_composition_score( out = {} u = vert.unique() for start, goal in regions: - end = verts_order.index(goal.get_next_poi(u)) + next_after_goal = goal.get_next_poi(u) + end = verts_order.index(next_after_goal) if next_after_goal is not None else verts_order.index(goal) + 1 labels = verts_order[verts_order.index(start) : end] vertebral_body = vert.extract_label(labels) * body_mask @@ -365,7 +374,7 @@ def body_composition_score( if name == "muscle": out[f"n_slices_{region_name}"] = n_slices - if height_m is not None and np.isfinite(mean_area): + if height_m is not None and height_m > 0 and np.isfinite(mean_area): out[f"muscle_index_{region_name}"] = round(mean_area / (height_m**2)) vat = out[f"mean_VAT_area_{region_name}"] From d0ebb89ac3bc2620f389a9d48f11ad5a580907aa Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Sun, 13 Sep 2026 09:19:24 +0000 Subject: [PATCH 05/26] speed up, split vert and ivd --- TPTBox/spine/spinestats/_run_all.py | 59 +++-- .../measure_ivd_and_vertebra_geometry.py | 210 ++++++++++++------ 2 files changed, 182 insertions(+), 87 deletions(-) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 9d01e402..d26f1bcf 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -260,7 +260,7 @@ def _need(*keys: str, compute: bool) -> bool: save_json(final_out, out) return out - logger.on_debug("load nii") + logger.on_debug("load nii", t2w_bf.get("sub")) t2w = to_nii(file_dict["t2w"]) if need_t2w or need_cobb else None vibe_water = to_nii(file_dict["vibe_part-water"], False) if need_vibe_wf else None vibe_fat = to_nii(file_dict["vibe_part-fat"], False) if need_vibe_wf else None @@ -317,7 +317,8 @@ def _need(*keys: str, compute: bool) -> bool: if need_bcs: logger.on_debug("body_composition_score") out["body_composition_score"] = body_composition_score(vibe_seg, vert, spine, dataset_id=100, height_m=height_m) - assert len(out["body_composition_score"]) != 0 + if len(out["body_composition_score"]) == 0: + logger.on_warning("body_composition_score returned empty (no vertebrae from configured regions present)") if need_mfi: logger.on_debug("muscle_fat_infiltration") out["muscle_fat_infiltration"] = muscle_fat_infiltration(vibe_water, vibe_fat, vibe_seg, vert, spine, roi=roi, dataset_id=100) @@ -410,30 +411,33 @@ def _flatten(prefix: str, obj: Any, out: dict[str, Any]) -> None: out[prefix] = obj -def _rows_from_json(subject_id: str, data: dict) -> tuple[dict[str, Any], list[dict[str, Any]]]: - """Split one subject's json into (per-subject row, per-vertebra rows). +def _rows_from_json(subject_id: str, data: dict) -> tuple[dict[str, Any], list[dict[str, Any]], list[dict[str, Any]]]: + """Split one subject's json into (per-subject row, per-vertebra rows, per-ivd rows). Per-subject row: everything except the per-label geometry dicts, flattened to dotted keys. - Per-vertebra rows: one row per label in ``ivd_geometry`` and - ``vert_geometry`` (source column indicates which). + Per-vertebra rows: one row per label in ``vert_geometry``. + Per-ivd rows: one row per label in ``ivd_geometry``. + Split by source so each output stays well below Excel's per-sheet + row limit (1_048_576). """ per_subject: dict[str, Any] = {"subject": subject_id} subject_view = {k: v for k, v in data.items() if k not in ("ivd_geometry", "vert_geometry")} _flatten("", subject_view, per_subject) per_vert: list[dict[str, Any]] = [] - for source_key in ("vert_geometry", "ivd_geometry"): + per_ivd: list[dict[str, Any]] = [] + for source_key, sink in (("vert_geometry", per_vert), ("ivd_geometry", per_ivd)): section = data.get(source_key) or {} if not isinstance(section, dict): continue for label, metrics in section.items(): if not isinstance(metrics, dict): continue - row: dict[str, Any] = {"subject": subject_id, "source": source_key, "label": label} + row: dict[str, Any] = {"subject": subject_id, "label": label} row.update(metrics) - per_vert.append(row) - return per_subject, per_vert + sink.append(row) + return per_subject, per_vert, per_ivd def _collector_worker( @@ -441,6 +445,7 @@ def _collector_worker( out_folder: Path, per_subject_name: str, per_vertebra_name: str, + per_ivd_name: str, flush_every: int, ) -> None: import pandas as pd # local import so the main process starts fast @@ -449,6 +454,7 @@ def _collector_worker( out_folder.mkdir(parents=True, exist_ok=True) subject_rows: list[dict[str, Any]] = [] vertebra_rows: list[dict[str, Any]] = [] + ivd_rows: list[dict[str, Any]] = [] seen: set[str] = set() def _flush() -> None: @@ -456,6 +462,8 @@ def _flush() -> None: pd.DataFrame(subject_rows).to_excel(out_folder / per_subject_name, index=False) if vertebra_rows: pd.DataFrame(vertebra_rows).to_excel(out_folder / per_vertebra_name, index=False) + if ivd_rows: + pd.DataFrame(ivd_rows).to_excel(out_folder / per_ivd_name, index=False) while True: try: @@ -472,9 +480,10 @@ def _flush() -> None: data = load_json(Path(json_path)) except Exception: continue - per_subj, per_vert = _rows_from_json(str(subject_id), data) + per_subj, per_vert, per_ivd = _rows_from_json(str(subject_id), data) subject_rows.append(per_subj) vertebra_rows.extend(per_vert) + ivd_rows.extend(per_ivd) seen.add(subject_id) if flush_every and len(seen) % flush_every == 0: _flush() @@ -499,11 +508,13 @@ def __init__( out_folder: str | Path, per_subject_name: str = "per_subject.xlsx", per_vertebra_name: str = "per_vertebra.xlsx", + per_ivd_name: str = "per_ivd.xlsx", flush_every: int = 200, ) -> None: self.out_folder = Path(out_folder) self.per_subject_name = per_subject_name self.per_vertebra_name = per_vertebra_name + self.per_ivd_name = per_ivd_name self.flush_every = flush_every self._queue: mp.Queue = mp.Queue() self._proc: mp.Process | None = None @@ -513,7 +524,14 @@ def start(self) -> None: return self._proc = mp.Process( target=_collector_worker, - args=(self._queue, self.out_folder, self.per_subject_name, self.per_vertebra_name, self.flush_every), + args=( + self._queue, + self.out_folder, + self.per_subject_name, + self.per_vertebra_name, + self.per_ivd_name, + self.flush_every, + ), daemon=True, ) self._proc.start() @@ -583,6 +601,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: if __name__ == "__main__": + import os from concurrent.futures import ProcessPoolExecutor, as_completed import pandas as pd @@ -590,10 +609,10 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: from TPTBox import No_Logger log = No_Logger() - + os.nice(20) OUT_FOLDER = Path("/DATA/NAS/ongoing_projects/robert/test/NAKO-stats") OUT_FOLDER.mkdir(parents=True, exist_ok=True) - N_CPUS = 10 # set >1 to parallelize + N_CPUS = 1 # set >1 to parallelize OVERRIDE = False aggregate = True do_not_update = False @@ -613,11 +632,13 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: else: # subjects = tqdm(loop_over_repaired_nako(test=False, sort=aggregate), total=30645) l = loop_over_repaired_nako(test=False, sort=aggregate) - total = 1000 - subjects = iter([next(l) for _ in range(total)]) + # total = 1000 + # subjects = iter([next(l) for _ in range(total)]) + subjects = l + # print(f"Run on {total=} random subset") if N_CPUS <= 1: - for f in subjects: + for f in tqdm(subjects, total=total): sub_id, missing, _ = _run_one((f, OVERRIDE, do_not_update)) if missing is not None: logger.on_fail("missing", list(f.keys()), missing) @@ -628,8 +649,8 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: else: from itertools import islice - with ProcessPoolExecutor(max_workers=N_CPUS) as ex: - batch_size = 1000 + with ProcessPoolExecutor(max_workers=N_CPUS, max_tasks_per_child=100) as ex: + batch_size = 100 l = tqdm(total=total) while True: gc.collect() diff --git a/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py index f16ce481..163a72ba 100644 --- a/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py +++ b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py @@ -177,20 +177,29 @@ def measure_ivd_and_vertebra_geometry( ) if instance_labels is None: instance_labels = [int(i) for i in vert.unique() if i > structure_label and i < structure_label + 100] + + # Precompute the spinal-canal reference signal once (was recomputed for every label). + if t2w is not None: + t2w_arr, spinal_canal_signal, spinal_canal_signal_old = _compute_spinal_canal_reference(t2w, vert, spine, erode=erode) + else: + t2w_arr = None + spinal_canal_signal = np.nan + spinal_canal_signal_old = np.nan + for label in instance_labels: if label == 26: continue try: raw = {} - # Isolate the current structure (disc or vertebra). - structure_mask = vert.extract_label(label) # 1. volume, central height, mean diameter info, raw = _compute_basic_geometry(vert, poi, label, step_size_mm=2, raw=raw, structure_label=structure_label) # 2. x1-x6 directional heights/widths - raw = _compute_directional_heights_widths(vert, structure_mask * spine, poi, label, step_size_mm=step_size_mm, raw=raw) + raw = _compute_directional_heights_widths(vert, poi, label, step_size_mm=step_size_mm, raw=raw) # 3. normalized T2 signal - if t2w is not None: - raw = _compute_t2_signal_ratio(t2w, vert, spine, label, raw=raw, erode=erode) + if t2w_arr is not None: + raw = _compute_t2_signal_ratio( + t2w_arr, vert, label, spinal_canal_signal, spinal_canal_signal_old, raw=raw, erode=erode + ) results[label] = _result_from_info(info) except Exception as e: results[label] = _nan_result(error=str(e)) @@ -440,6 +449,50 @@ def _swap(a, b): return b, a +def _batched_ray_segments(mesh: trimesh.Trimesh, ray_direction: np.ndarray, origins: np.ndarray): + """Cast a batch of parallel rays through ``mesh`` and return per-ray segment lengths and endpoints. + + trimesh's ``intersects_location`` accepts arrays of origins/directions and processes them in one + C call, which is dramatically faster than looping in Python. Rays that miss return length 0 and + zero-vector endpoints. Endpoints are the first and last locations trimesh returns for each ray + (matches the single-ray behaviour of :func:`_segment_length`). + """ + n = origins.shape[0] + directions = np.broadcast_to(ray_direction, (n, 3)) + locations, index_ray, _ = mesh.ray.intersects_location( + ray_origins=origins, ray_directions=directions, multiple_hits=True + ) + lengths = np.zeros(n) + first_pts = np.zeros((n, 3)) + last_pts = np.zeros((n, 3)) + if len(index_ray) == 0: + return lengths, first_pts, last_pts + order = np.argsort(index_ray, kind="stable") + idx_sorted = index_ray[order] + loc_sorted = locations[order] + change = np.empty(len(idx_sorted), dtype=bool) + change[0] = True + change[1:] = idx_sorted[1:] != idx_sorted[:-1] + first_pos = np.flatnonzero(change) + last_pos = np.empty_like(first_pos) + last_pos[:-1] = first_pos[1:] - 1 + last_pos[-1] = len(idx_sorted) - 1 + ray_ids = idx_sorted[first_pos] + first_pts[ray_ids] = loc_sorted[first_pos] + last_pts[ray_ids] = loc_sorted[last_pos] + lengths[ray_ids] = np.linalg.norm(first_pts[ray_ids] - last_pts[ray_ids], axis=1) + return lengths, first_pts, last_pts + + +def _grid_origins(base: np.ndarray, v1: np.ndarray, v2: np.ndarray, xs: np.ndarray, ys: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Return (origins, flat_x, flat_y) for a 2D grid of ray origins in the (v1, v2) plane.""" + grid_x, grid_y = np.meshgrid(xs, ys, indexing="ij") + flat_x = grid_x.ravel() + flat_y = grid_y.ravel() + origins = base[None, :] + flat_x[:, None] * v1[None, :] + flat_y[:, None] * v2[None, :] + return origins, flat_x, flat_y + + # --------------------------------------------------------------------------- # Mesh + orientation extraction (shared between IVD and vertebra mode) # --------------------------------------------------------------------------- @@ -469,8 +522,19 @@ def _get_mesh_and_directions(nii: NII, poi: POI | None, label: int, raw: dict, r (surface mesh, (up, dir_1, dir_2) direction vectors, binary segmentation) """ voxel_size = nii.zoom - nii = nii.apply_crop(nii.compute_crop(dist=1)) - segmentation = nii.extract_label(label) + # Whole-spine crop (kept as-is): tightening to the label's own bbox changes the array shape + # passed to marching_cubes, which changes the vertex/face emission order. trimesh's ray + # backend uses that order to return intersection locations, and _segment_length picks + # ``locs[0]`` and ``locs[-1]`` -- so a different order silently changes x1-x6 for + # non-convex intersections (mostly vertebrae). Reuse the crop across labels via ``raw``. + nii_cropped = raw.get("nii_cropped") + if nii_cropped is None: + nii_cropped = nii.apply_crop(nii.compute_crop(dist=1)) + raw["nii_cropped"] = nii_cropped + segmentation = raw.get("label_mask") + if segmentation is None: + segmentation = nii_cropped.extract_label(label) + raw["label_mask"] = segmentation arr = segmentation.get_array() if label < 100: @@ -483,13 +547,11 @@ def _get_mesh_and_directions(nii: NII, poi: POI | None, label: int, raw: dict, r post /= norm(post) right /= norm(right) if label == 2: - up = _pca_principal_axes(arr, voxel_size, up_axis=nii.get_axis("S"))[0] + up = _pca_principal_axes(arr, voxel_size, up_axis=nii_cropped.get_axis("S"))[0] direction_vectors = (up, post, right) else: - up_axis = nii.get_axis("S") if label < 100 else -1 - direction_vectors = ( - raw["direction_vectors"] if "direction_vectors" in raw else _pca_principal_axes(arr, voxel_size, up_axis=up_axis) - ) + cached = raw.get("direction_vectors") + direction_vectors = cached if cached is not None else _pca_principal_axes(arr, voxel_size, up_axis=-1) mesh = raw["mesh"] if "mesh" in raw and not recompute_mesh else _segmentation_to_surface_mesh(arr, voxel_size) raw["mesh"] = mesh @@ -575,18 +637,17 @@ def _compute_basic_geometry( # ------------------------------------------------------------------ if "list_heights" not in raw or len(raw["list_heights"]) == 0: mesh, direction_vectors, _ = _get_mesh_and_directions(nii, poi, label, raw) - + up_vector, v1, v2 = direction_vectors search_radius = int(info.mean_diameter) - sampled_heights = [] - - for x in range(-search_radius, search_radius, step_size_mm): - for y in range(-search_radius, search_radius, step_size_mm): - intersections = _intersect_ray_from_dirs(direction_vectors, mesh, x, y) - height = _segment_length(intersections) - if height > 0: - sampled_heights.append(height) - - raw["list_heights"] = sampled_heights + if search_radius > 0: + xs = np.arange(-search_radius, search_radius, step_size_mm) + ys = np.arange(-search_radius, search_radius, step_size_mm) + base = mesh.center_mass - up_vector * 1000 + origins, _, _ = _grid_origins(base, v1, v2, xs, ys) + lengths, _, _ = _batched_ray_segments(mesh, up_vector, origins) + raw["list_heights"] = lengths[lengths > 0].tolist() + else: + raw["list_heights"] = [] info._update_height_statistics(raw["list_heights"]) @@ -600,17 +661,15 @@ def _local_max_height(mesh, direction_vectors, around_point, search_diameter, st mesh edge and under-measure; sampling a small patch around the point and taking the max is more robust. """ - sampled_heights = [] + up_vector, v1, v2 = direction_vectors d = int(search_diameter / step_size_mm) - for xi in range(-d, d): - x = xi * step_size_mm - for yi in range(-d, d): - y = yi * step_size_mm - intersections = _intersect_ray_from_dirs(direction_vectors, mesh, x, y, around_point) - height = _segment_length(intersections) - if height != 0: - sampled_heights.append(height) - return np.max(sampled_heights) + xs = np.arange(-d, d) * step_size_mm + ys = np.arange(-d, d) * step_size_mm + base = np.asarray(around_point) - up_vector * 1000 + origins, _, _ = _grid_origins(base, v1, v2, xs, ys) + lengths, _, _ = _batched_ray_segments(mesh, up_vector, origins) + lengths = lengths[lengths > 0] + return float(np.max(lengths)) def _max_diameter_in_plane(ray_vector, v1, v2, mesh, diameter: float = 30, step_size_mm: float = 2.0): @@ -629,25 +688,18 @@ def _max_diameter_in_plane(ray_vector, v1, v2, mesh, diameter: float = 30, step_ (max_width, point_1, point_2, grid_x, grid_y) for the widest ray found. """ d = ceil(diameter / step_size_mm) - best_width = 0 - out = (0.0, None, None, 0.0, 0.0) - for xi in range(-d, d): - x = xi * step_size_mm - for yi in range(-d, d): - y = yi * step_size_mm - intersections = _intersect_ray(ray_vector, v1, v2, mesh, x, y) - width = _segment_length(intersections) - if width == 0: - continue - p1 = intersections[0] - p2 = intersections[-1] - if best_width < width: - best_width = width - out = (width, p1.round(2), p2.round(2), x, y) - return out - - -def _compute_directional_heights_widths(nii: NII, subreg: NII, poi, label: int = 123, step_size_mm: float = 0.5, raw: dict | None = None): # noqa: ARG001 + xs = np.arange(-d, d) * step_size_mm + ys = np.arange(-d, d) * step_size_mm + base = mesh.center_mass - ray_vector * 1000 + origins, flat_x, flat_y = _grid_origins(base, v1, v2, xs, ys) + lengths, first_pts, last_pts = _batched_ray_segments(mesh, ray_vector, origins) + if lengths.size == 0 or lengths.max() == 0: + return 0.0, None, None, 0.0, 0.0 + k = int(np.argmax(lengths)) + return float(lengths[k]), first_pts[k].round(2), last_pts[k].round(2), float(flat_x[k]), float(flat_y[k]) + + +def _compute_directional_heights_widths(nii: NII, poi, label: int = 123, step_size_mm: float = 0.5, raw: dict | None = None): """Compute the x1-x6 directional heights and widths for one structure (stage 2). How it's computed @@ -675,7 +727,8 @@ def _compute_directional_heights_widths(nii: NII, subreg: NII, poi, label: int = if info.x_values: return raw try: - mesh, direction_vectors, _ = _get_mesh_and_directions(nii, poi, label, raw, recompute_mesh=True) + # Reuse the mesh built in stage 1 (previously rebuilt here for no reason). + mesh, direction_vectors, _ = _get_mesh_and_directions(nii, poi, label, raw) up = direction_vectors[0] center = np.asarray(poi[label % 100, Location.Vertebra_Corpus], dtype=float) @@ -724,13 +777,43 @@ def _compute_directional_heights_widths(nii: NII, subreg: NII, poi, label: int = return raw -def _compute_t2_signal_ratio( +def _compute_spinal_canal_reference( t2w_nii: NII, nii: NII, subregs: NII, - label: int = 123, + erode: int = 1, + spinal_bins: int = 64, + spinal_peak_frac_height: float = 0.5, +) -> tuple[np.ndarray, float, float]: + """Prepare the shared T2 array and spinal-canal reference signals once per subject. + + The spinal canal (subregion 61) is the same for every structure, so + eroding it and reducing its intensity to a scalar used to be repeated + for every label; now it is done once and reused. + + Returns: + ------- + tuple[np.ndarray, float, float] + (t2w_array, spinal_canal_signal_peak, spinal_canal_signal_mean) + """ + if t2w_nii.shape != nii.shape: + t2w_nii.resample_from_to_(nii, verbose=False) + spinal_mask = subregs.extract_label(61).erode_msk(erode, connectivity=1, verbose=False) + t2w_arr = t2w_nii.get_array() + spinal_vals = t2w_arr[spinal_mask.get_array().astype(bool)] + spinal_canal_signal_old = float(np.mean(spinal_vals)) if spinal_vals.size > 0 else np.nan + spinal_canal_signal = peak_centered_mean(spinal_vals, bins=spinal_bins, peak_frac_height=spinal_peak_frac_height) + return t2w_arr, spinal_canal_signal, spinal_canal_signal_old + + +def _compute_t2_signal_ratio( + t2w_arr: np.ndarray, + nii: NII, + label: int, + spinal_canal_signal: float, + spinal_canal_signal_old: float, raw: dict | None = None, - erode=1, + erode: int = 1, spinal_bins: int = 64, spinal_peak_frac_height: float = 0.5, ): @@ -739,9 +822,9 @@ def _compute_t2_signal_ratio( How it's computed ------------------ The mean T2 intensity inside the (slightly eroded, to avoid partial-volume - edge voxels) structure mask is divided by the mean T2 intensity in the - spinal canal (subregion label 61, also eroded). The spinal canal is used - as an internal reference to normalize away scanner/sequence-dependent + edge voxels) structure mask is divided by the precomputed mean T2 intensity + in the spinal canal (subregion label 61, also eroded). The spinal canal is + used as an internal reference to normalize away scanner/sequence-dependent intensity scaling. The spinal canal segmentation may contain darker structures such as @@ -765,22 +848,13 @@ def _compute_t2_signal_ratio( info: _StructureMeasurements = raw["info"] if info.signal_values: return raw - if t2w_nii.shape != nii.shape: - t2w_nii.resample_from_to_(nii, verbose=False) structure_mask = nii.extract_label(label) eroded_mask = structure_mask.erode_msk(erode, connectivity=1, verbose=False) structure_mask = eroded_mask if eroded_mask.sum() != 0 else structure_mask - spinal_mask = subregs.extract_label(61).erode_msk(erode, connectivity=1, verbose=False) - - t2w_arr = t2w_nii.get_array() structure_vals = t2w_arr[structure_mask.get_array().astype(bool)] - spinal_vals = t2w_arr[spinal_mask.get_array().astype(bool)] structure_signal_old = float(np.mean(structure_vals)) if structure_vals.size > 0 else np.nan - spinal_canal_signal_old = float(np.mean(spinal_vals)) if spinal_vals.size > 0 else np.nan - structure_signal = peak_centered_mean(structure_vals, bins=spinal_bins, peak_frac_height=spinal_peak_frac_height) - spinal_canal_signal = peak_centered_mean(spinal_vals, bins=spinal_bins, peak_frac_height=spinal_peak_frac_height) info.signal = structure_signal / spinal_canal_signal info.structure_signal = structure_signal From b12cf446783eab04e3168cce94855269fd5be3a0 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Sun, 13 Sep 2026 09:19:36 +0000 Subject: [PATCH 06/26] update Readme --- TPTBox/spine/spinestats/README.md | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/TPTBox/spine/spinestats/README.md b/TPTBox/spine/spinestats/README.md index 7e04fed7..3282494e 100644 --- a/TPTBox/spine/spinestats/README.md +++ b/TPTBox/spine/spinestats/README.md @@ -333,15 +333,22 @@ Implementation notes: ## Excel collector `ExcelCollector` in `_run_all.py` runs a background process that turns -each finished json into two rolling Excel files in a configurable +each finished json into three rolling Excel files in a configurable folder: - `per_subject.xlsx` — one row per subject with every scalar top-level metric flattened to dotted keys (e.g. `VBQ_score.VBQ_L1-L4`, `torso_vat_sat_muscle_mass.VAT`). -- `per_vertebra.xlsx` — one row per (subject, label), populated from - `vert_geometry` and `ivd_geometry`. The `source` column indicates - which of the two sections the row came from. + `ivd_geometry` and `vert_geometry` are excluded here. +- `per_vertebra.xlsx` — one row per (subject, label) from + `vert_geometry` (vertebra bodies). +- `per_ivd.xlsx` — one row per (subject, label) from `ivd_geometry` + (intervertebral discs). + +The vertebra and IVD tables were split so that the full NAKO cohort +stays under Excel's per-sheet row limit (1 048 576 rows). A single +combined table would exceed that once the cohort passes ~23 k subjects +with ~23 labels per section. Usage: From 69fdf16ba1da50837408dc0ef394ada239a22802 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Sun, 13 Sep 2026 09:20:08 +0000 Subject: [PATCH 07/26] update nako loader --- TPTBox/spine/spinestats/_load_nako.py | 8 ++++++-- TPTBox/spine/spinestats/all_output_reference.md | 16 +++++++++++----- 2 files changed, 17 insertions(+), 7 deletions(-) diff --git a/TPTBox/spine/spinestats/_load_nako.py b/TPTBox/spine/spinestats/_load_nako.py index 9b77f311..38e1a8d3 100644 --- a/TPTBox/spine/spinestats/_load_nako.py +++ b/TPTBox/spine/spinestats/_load_nako.py @@ -234,8 +234,12 @@ def loop_over_repaired_nako( if "PatientSize" in js: subj_dict["height_m"] = js["PatientSize"] break - except json.decoder.JSONDecodeError: - log.on_fail(f, "json.decoder.JSONDecodeError") + except json.decoder.JSONDecodeError as e: + json_path = f.file.get("json", f) + log.on_fail( + f"json.decoder.JSONDecodeError while reading {json_path}: {e} " + f"(subject={sub}, dataset={dataset}); continuing with next sidecar" + ) if verbose: log.on_log(sub) mapping = {"T2w": "t2w"} diff --git a/TPTBox/spine/spinestats/all_output_reference.md b/TPTBox/spine/spinestats/all_output_reference.md index acff5291..355a2e32 100644 --- a/TPTBox/spine/spinestats/all_output_reference.md +++ b/TPTBox/spine/spinestats/all_output_reference.md @@ -248,15 +248,21 @@ Implementation notes: ## Excel collector -`ExcelCollector` in `all.py` runs a background process that turns each -finished json into two rolling Excel files in a configurable folder: +`ExcelCollector` in `_run_all.py` runs a background process that turns +each finished json into three rolling Excel files in a configurable +folder: - `per_subject.xlsx` — one row per subject with every scalar top-level metric flattened to dotted keys (e.g. `VBQ_score.VBQ_L1-L4`, `torso_vat_sat_muscle_mass.VAT`). -- `per_vertebra.xlsx` — one row per (subject, label), populated from - `vert_geometry` and `ivd_geometry`. The `source` column indicates - which of the two sections the row came from. + `ivd_geometry` and `vert_geometry` are excluded here. +- `per_vertebra.xlsx` — one row per (subject, label) from + `vert_geometry` (vertebra bodies). +- `per_ivd.xlsx` — one row per (subject, label) from `ivd_geometry` + (intervertebral discs). + +The vertebra and IVD tables are split so that the full NAKO cohort stays +under Excel's per-sheet row limit (1 048 576 rows). Usage: From 4bb671499ee56d8be9be2efdd450973c087af164 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Sun, 13 Sep 2026 09:20:20 +0000 Subject: [PATCH 08/26] nako cannonical --- TPTBox/spine/spinestats/_load_nako_wh.py | 630 +++++++++++++++++++++-- 1 file changed, 578 insertions(+), 52 deletions(-) diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index 2e819cbc..c076ffdb 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -1,5 +1,6 @@ import json import os +import tempfile from concurrent.futures import ProcessPoolExecutor, as_completed from pathlib import Path @@ -18,6 +19,28 @@ _DEFAULT_DECISION_CACHE = Path(__file__).with_name("_load_nako_wh_decisions.json") +_NON_INTERACTIVE = False + + +class _non_interactive_mode: + """Context manager that flips the module-level ``_NON_INTERACTIVE`` flag. + + While active, every prompt site in this module short-circuits to the same + behaviour it would use when the user typed "skip" — first candidate is kept + and nothing is written to the decision cache. + """ + + def __enter__(self): + global _NON_INTERACTIVE # noqa: PLW0603 + self._prev = _NON_INTERACTIVE + _NON_INTERACTIVE = True + return self + + def __exit__(self, *exc): + global _NON_INTERACTIVE # noqa: PLW0603 + _NON_INTERACTIVE = self._prev + return False + class DecisionCache: """Persistent per-(subject, key) decision cache backed by a JSON file. @@ -82,6 +105,8 @@ def _prompt_choice(sub: str, key: str, question: str, options: list[str], allow_ - or the sentinel string ``"__skip__"`` (do not save; ask again next run). ``reason`` is the free-text reason string, or ``None`` when skipped. """ + if _NON_INTERACTIVE: + return "__skip__", None print("\n" + "=" * 72) print(f"[decision needed] subject={sub} key={key}") print(question) @@ -123,6 +148,12 @@ def resolve_pick( if cached_pick in labels: return candidates[labels.index(cached_pick)] log.on_warning(f"cached decision {cached_pick!r} for ({sub},{key}) no longer matches candidates; re-asking") + if key == "main:pd" or key.endswith(":pd"): + # auto-accept: for duplicate PDs always prefer the higher ID (lexicographically last label). + # Covers both the initial pick (``main:pd``) and the cross-family conflict path + # (``main-conflict:pd``, ``mevibe-conflict:pd``, ``vibe-conflict:pd``, …). + idx = max(range(len(labels)), key=labels.__getitem__) + return candidates[idx] # transient: do not save choice, reason = _prompt_choice(sub, key, question, labels, allow_discard=allow_discard) if choice == "__skip__": return candidates[0] if candidates else None # transient: do not save @@ -172,10 +203,11 @@ def verify_missing_images(cache: DecisionCache, sub: str, subj_dict: dict) -> No cached = cache.get(sub, key) decision = cached.get("decision") if isinstance(cached, dict) else cached Print_Logger().on_debug(sub, key, cached) - if key in ["missing:T2haste", "missing:vibe_part-inphase", "missing:eco0-opp1", "missing:T2w"]: + if decision is None and _NON_INTERACTIVE: + decision = "keep" # transient: skip prompt, do not save + elif key in ["missing:T2haste", "missing:vibe_part-inphase", "missing:eco0-opp1", "missing:T2w"]: decision = "missing" cache.set(sub, key, {"decision": decision, "reason": "Missing"}) - elif decision is None: print("\n" + "=" * 72) print(f"[decision needed] subject={sub} base image {base!r} not found.") @@ -201,36 +233,135 @@ def verify_missing_images(cache: DecisionCache, sub: str, subj_dict: dict) -> No subj_dict[dep] = None -def check_same_grid(cache: DecisionCache, sub: str, group: str, files: list) -> bool: - """Verify all ``files`` share the same grid (spacing/shape/affine via ``bf.get_grid_info()``). +_RESAMPLE_TMP_ROOT = Path( + os.environ.get( + "TPTBOX_RESAMPLED_SCRATCH", + "/DATA/NAS/datasets_processed/NAKO/_resampled_scratch", + ) +) + + +def _resample_to_ref(bf, ref_nii): + """Resample ``bf`` onto ``ref_nii``'s grid and persist under a scratch root. - On mismatch: print spacing per file, log a warning, and record an "issue" entry in the cache - (once per subject/group/grid-signature) so the mismatch is surfaced but not re-prompted. - Returns True when grids match, False otherwise. + Files are written to ``$TPTBOX_RESAMPLED_SCRATCH`` (default + ``/DATA/NAS/datasets_processed/NAKO/_resampled_scratch``) mirroring the + original BIDS layout with a ``desc-resampled`` entity added, so we neither + dirty the read-only NAKO source tree nor cross filesystems (the scratch + root lives on the same disk as NAKO so :func:`os.link` still works when + :func:`hard_link` later re-links the file into the canonical dataset). + The file is only written when it does not already exist; the returned + :class:`BIDS_FILE` always points at the resampled path. + """ + existing_desc = bf.get("desc", None) + new_desc = "resampled" if not existing_desc else f"{existing_desc}Resampled" + _RESAMPLE_TMP_ROOT.mkdir(parents=True, exist_ok=True) + new_bf = bf.get_changed_bids( + file_type="nii.gz", + bids_format=bf.bids_format, + parent=bf.parent, + info={"desc": new_desc}, + dataset_path=str(_RESAMPLE_TMP_ROOT), + ) + new_path = new_bf.file["nii.gz"] + if not new_path.exists(): + new_path.parent.mkdir(parents=True, exist_ok=True) + src_nii = to_nii(bf) + resampled = src_nii.resample_from_to(ref_nii) + resampled.save(new_path) + return new_bf + + +def check_same_grid(cache: DecisionCache, sub: str, group: str, files_by_key: dict, inphase_key: str) -> bool: + """Ensure every file in ``files_by_key`` shares the grid of ``files_by_key[inphase_key]``. + + On mismatch the user is prompted (with a persistent per-(sub, group) decision): + * ``y`` → resample each mismatched file to the inphase grid, save it beside + the original with an added ``desc-resampled`` entity, and update + ``files_by_key[k]`` in place to point at the resampled file. + * ``n`` → drop the mismatched keys (``files_by_key[k] = None``). + + Under ``_NON_INTERACTIVE`` the default is ``resample`` and the decision is not + persisted. Returns ``True`` when every entry ended up on the reference grid. """ grids: dict = {} - for bf in files: + for k, bf in files_by_key.items(): if bf is None: continue + # Segmentations (msk) don't participate in the image-grid check. + if getattr(bf, "format", None) == "msk" or k.startswith("msk"): + continue + # Skip BIDS entries that carry no NIfTI (e.g. POI files with only .json / .mrk.json). + get_nii_file = getattr(bf, "get_nii_file", None) + if callable(get_nii_file) and get_nii_file() is None: + continue try: g = bf.get_grid_info() except Exception as e: # noqa: BLE001 log.on_warning(f"get_grid_info failed for {_fmt_file(bf)}: {e}") + files_by_key[k] = None continue - grids.setdefault(str(g), []).append(_fmt_file(bf)) - if len(grids) <= 1: + grids[k] = (g, str(g)) + + unique_sigs = {sig for _, sig in grids.values()} + if len(unique_sigs) <= 1: + return True + + if inphase_key not in grids: + log.on_warning(f"inphase reference {inphase_key!r} missing for subject {sub}; cannot resample") + return False + + ref_sig = grids[inphase_key][1] + mismatched = [k for k, (_, sig) in grids.items() if sig != ref_sig] + if not mismatched: return True - key = f"grid_mismatch:{group}:" + "|".join(sorted(grids.keys())) + + if _NON_INTERACTIVE: + # Precompute / non-interactive mode: don't touch anything, don't resample. + # Grid info for every file was already read (populating the JSON cache), which + # is all --precompute-grid wants; leave the actual reconciliation to a later + # interactive run. + log.on_warning(f"grid mismatch in {group} for subject {sub} (skipped in non-interactive mode)") + return False + print("\n" + "=" * 72) - print(f"[ISSUE] subject={sub} group={group} grid mismatch across {sum(len(v) for v in grids.values())} files:") - for g, names in grids.items(): - print(f" grid {g}") - for n in names: - print(f" - {n}") - log.on_warning(f"grid mismatch in {group} for subject {sub}") - # if cache.get(sub, key) is None: TODO - # cache.set(sub, key, {"issue": "grid_mismatch", "grids": {g: n for g, n in grids.items()}}) - return False + print(f"[grid mismatch] subject={sub} group={group} reference={inphase_key} ({ref_sig})") + for k in mismatched: + print(f" {k}: {grids[k][1]} -> {_fmt_file(files_by_key[k])}") + + cache_key = f"grid_mismatch:{group}" + cached = cache.get(sub, cache_key) + decision = cached.get("decision") if isinstance(cached, dict) else cached + if decision not in ("resample", "remove"): + print(f" [y] resample mismatched files to the {inphase_key} grid") + print(" [n] drop the mismatched keys") + print(" [s] skip (do not save; ask again next run)") + while True: + raw = input("> ").strip().lower() + if raw in ("y", "n", "s"): + break + print("invalid input, try again") + if raw == "s": + decision = "resample" # transient + else: + decision = "resample" if raw == "y" else "remove" + reason = "resample-to-inphase" if raw == "y" else _prompt_reason() + cache.set(sub, cache_key, {"decision": decision, "reason": reason}) + + if decision == "remove": + for k in mismatched: + files_by_key[k] = None + return False + + ref_nii = to_nii(files_by_key[inphase_key]) + for k in mismatched: + src_bf = files_by_key[k] + try: + files_by_key[k] = _resample_to_ref(src_bf, ref_nii) + except Exception as e: # noqa: BLE001 + log.on_warning(f"resample failed for {_fmt_file(src_bf)}: {e}; dropping key {k!r}") + files_by_key[k] = None + return True def _cached_pick(cached): @@ -259,6 +390,8 @@ def resolve_keep_chunks( if cached is not None: drop = cached.get("drop") if isinstance(cached, dict) else cached return list(drop) if isinstance(drop, list) else [] + if _NON_INTERACTIVE: + return [] # transient: keep all, do not save print("\n" + "=" * 72) print(f"[decision needed] subject={sub} unknown t2w chunks (not in BWS/LWS/HWS)") for c in unknown_chunks: @@ -434,8 +567,10 @@ def loop_over_repaired_nako( test=False, verbose=False, sort=True, - test_key="/110/110", # path matching. if you want on specific us a 6 digits + test_key="/100/10", # path matching. if you want on specific us a 6 digits decision_cache: DecisionCache | Path | str | None = None, + corrected_index: dict | Path | str | None = None, + skip_subject=None, ): """Iterate over the repaired NAKO dataset yielding per-subject file dicts. @@ -466,19 +601,28 @@ def loop_over_repaired_nako( decision_cache = DecisionCache(decision_cache) if decision_cache is not None else DecisionCache() cache = decision_cache + if isinstance(corrected_index, (str, Path)): + corrected_index = load_corrected_index(Path(corrected_index)) + elif corrected_index is None: + corrected_index = {} + gbi = Buffered_BIDS_Global_info( datasets=dataset, parents=[ "rawdata", "rawdata_stitched", "derivatives_Abdominal-Segmentation", - # "derivatives_mevibe", #copied into "derivatives_Abdominal-Segmentation" + "derivatives_mevibe", # partial overlap with Abdominal-Segmentation, but hosts the paraspinal-muscles reconstructed masks "derivatives_inversion", + "derivatives-fullbody-poi", # fullbody / fov101 / fov102 POIs + registered segmentations on the stitched-water grid ], filter_file=(lambda x: test_key in str(x)) if test else None, + ) for sub, subj in gbi.enumerate_subjects(sort=sort, shuffle=not sort): + if skip_subject is not None and skip_subject(sub): + continue subj_dict = {"id": sub, "dataset": dataset} # Primary source: baseline CSV, height is in cm. if verbose: @@ -543,7 +687,16 @@ def loop_over_repaired_nako( else: subj_dict[k] = picked - keys = ["msk_seg-body-composition_mod-mevibe"] + # Extra mevibe segmentations sourced from `derivatives_mevibe`. When the same + # filename also lives in `derivatives_Abdominal-Segmentation` we auto-prefer + # the `derivatives_mevibe` copy below to avoid an interactive resolve_pick. + _MEVIBE_EXTRA_KEYS = ( + "msk_seg-spine_mod-mevibe_part-eco0-opp1", + "msk_seg-vert_mod-mevibe_part-eco0-opp1", + "msk_seg-seg-paraspinal-muscles-517-post_mod-mevibe_part-fat-fraction_desc-reconstructed-percent-20", + "msk_seg-seg-paraspinal-muscles-517-post-figure_mod-mevibe_part-fat-fraction_desc-reconstructed-percent-20", + ) + keys = ["msk_seg-body-composition_mod-mevibe", *_MEVIBE_EXTRA_KEYS] if add_mevibe: q = subj.new_query() q.filter_format("mevibe") @@ -557,21 +710,21 @@ def loop_over_repaired_nako( elif cached_pick in labels: mevibe_fams = [mevibe_fams[labels.index(cached_pick)]] else: - choice, reason = _prompt_choice(sub, "mevibe_fam", "Multiple mevibe families; pick one.", labels, allow_discard=True) - if choice == "__skip__": - mevibe_fams = mevibe_fams[:1] - elif choice is None: - cache.set(sub, "mevibe_fam", {"pick": "__discard__", "reason": reason}) - mevibe_fams = [] - else: - cache.set(sub, "mevibe_fam", {"pick": labels[choice], "reason": reason}) - mevibe_fams = [mevibe_fams[choice]] + # auto-accept: always prefer the higher-ID mevibe cluster (lexicographically last label). + idx = max(range(len(labels)), key=labels.__getitem__) + mevibe_fams = [mevibe_fams[idx]] # transient: do not save for fam in mevibe_fams: mevibe_out = get_corrected_mevibe(fam, compute_PDFF=compute_PDFF) - check_same_grid(cache, sub, "mevibe", list(mevibe_out.values())) + check_same_grid(cache, sub, "mevibe", mevibe_out, inphase_key="eco3-in1") subj_dict = {**mevibe_out, **subj_dict} for k, v in fam.items(): if k in keys: + # For the seg files also present under `derivatives_Abdominal-Segmentation`, + # auto-prefer the `derivatives_mevibe` copy (avoids an interactive prompt). + if k in _MEVIBE_EXTRA_KEYS and len(v) > 1: + preferred = [bf for bf in v if "/derivatives_mevibe/" in _fmt_file(bf)] + if preferred: + v = preferred k = mapping.get(k, k) # noqa: PLW2901 if len(v) > 1: picked = resolve_pick(cache, sub, f"mevibe:{k}", f"Multiple mevibe files for {k}; pick one.", v) @@ -599,6 +752,18 @@ def loop_over_repaired_nako( q.filter_format("vibe") q.filter("chunk", lambda _: False, required=False) # q.filter("run", lambda x: x != "2", required=False) + # Fullbody-POI derivatives live on the stitched-water grid; they land in + # the same vibe family via shared (sub, sequ-stitched, acq, part-water) entities. + _FULLBODY_POI_KEYS = ( + "poi_seg-fullbody_part-water", + # "poi_seg-fov101-reg_part-water", + # "poi_seg-fov102-reg_part-water", + # "msk_seg-fov101-reg_part-water", + # "msk_seg-fov102-reg_part-water", + # "msk_seg-fov101-reg-split-seg_part-water", + # "msk_seg-fov102-reg-split-seg_part-water", + # "msk_seg-fov102-reg-split-seg-leg_part-water", + ) keys = [ "vibe_part-inphase", "vibe_part-outphase", @@ -606,6 +771,7 @@ def loop_over_repaired_nako( "vibe_part-water", "msk_seg-body-composition_mod-vibe", *mapping.keys(), + *_FULLBODY_POI_KEYS, ] vibe_fams = list(q.loop_dict(key_addendum=["mod", "part", "desc"])) if len(vibe_fams) > 1: @@ -626,7 +792,7 @@ def loop_over_repaired_nako( cache.set(sub, "vibe_fam", {"pick": labels[choice], "reason": reason}) vibe_fams = [vibe_fams[choice]] for fam in vibe_fams: - vibe_files = [] + vibe_by_key: dict = {} for _k in ( "vibe_part-inphase", "vibe_part-outphase", @@ -635,9 +801,16 @@ def loop_over_repaired_nako( "vibe_part-water_desc-reconstructed", "vibe_part-fat_desc-reconstructed", ): - if _k in fam: - vibe_files.extend(fam[_k]) - check_same_grid(cache, sub, "vibe", vibe_files) + if fam.get(_k): + vibe_by_key[_k] = fam[_k][0] + check_same_grid(cache, sub, "vibe", vibe_by_key, inphase_key="vibe_part-inphase") + # Propagate the check's outcome back to ``fam`` so the downstream unpack + # loop below picks up resampled files (or skips removed keys). + for _k, bf in vibe_by_key.items(): + if bf is None: + fam.data_dict.pop(_k, None) + else: + fam[_k] = [bf] for k, v in fam.items(): if k in keys: k = mapping.get(k, k) # noqa: PLW2901 @@ -665,6 +838,8 @@ def loop_over_repaired_nako( subj_dict["vert"] = vert subj_dict["spine"] = spine subj_dict["poi"] = poi + if corrected_index: + _apply_corrections_to_subj_dict(str(sub), subj_dict, corrected_index) verify_missing_images(cache, sub, subj_dict) yield subj_dict @@ -672,6 +847,33 @@ def loop_over_repaired_nako( allowed_keys = ["sub", "sequ", "ses", "seg", "acq", "chunk", "part", "mod", "desc", "rec"] +def _is_grid_only_json(path: Path | str) -> bool: + """True when ``path`` is a JSON sidecar whose only key is ``"grid"``. + + Such a file was written by :func:`_add_grid_info_to_json` on a sidecar that + did not exist before — it carries no real metadata, only cached grid info + for the associated NIfTI. We do not want these propagating into the + canonical dataset alongside the .nii.gz. + """ + try: + content = json.loads(Path(path).read_text()) + except (OSError, json.JSONDecodeError): + return False + return isinstance(content, dict) and set(content.keys()) == {"grid"} + + +_CANONICAL_DONE_ROOT = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical/.hardlink_done") + + +def _hard_link_done_marker(sub: str) -> Path: + return _CANONICAL_DONE_ROOT / f"{sub}.done" + + +def is_hard_linked(sub: str) -> bool: + """Return True when :func:`hard_link` has completed successfully for ``sub``.""" + return _hard_link_done_marker(str(sub)).exists() + + def hard_link( d: dict, dataset="/DATA/NAS/datasets_processed/NAKO/dataset-nako/", @@ -699,10 +901,24 @@ def hard_link( segs = [ "msk_seg-body-composition_mod-mevibe", # MEVIBE + # MEVIBE extras from `derivatives_mevibe` + "msk_seg-spine_mod-mevibe_part-eco0-opp1", + "msk_seg-vert_mod-mevibe_part-eco0-opp1", + "msk_seg-seg-paraspinal-muscles-517-post_mod-mevibe_part-fat-fraction_desc-reconstructed-percent-20", + "msk_seg-seg-paraspinal-muscles-517-post-figure_mod-mevibe_part-fat-fraction_desc-reconstructed-percent-20", "vibeseg100", # vibe "MRSegmentator", # vibe "msk_seg-body-composition_mod-vibe", # vibe "roi", # vibe + # Fullbody-POI derivatives (stitched-water grid) + "poi_seg-fullbody_part-water", + "poi_seg-fov101-reg_part-water", + "poi_seg-fov102-reg_part-water", + "msk_seg-fov101-reg_part-water", + "msk_seg-fov102-reg_part-water", + "msk_seg-fov101-reg-split-seg_part-water", + "msk_seg-fov102-reg-split-seg_part-water", + "msk_seg-fov102-reg-split-seg-leg_part-water", "vert", # t2w (stiched) "spine", # t2w (stiched) "poi", # t2w (stiched) @@ -726,7 +942,7 @@ def hard_link( for keys, parent in [(imgs, "rawdata"), (segs, "derivatives")]: for key in keys: bf = d.pop(key, None) - if bf is None: + if bf is None or bf == "": continue info = {"run": None} if isinstance(bf, str): @@ -742,12 +958,306 @@ def hard_link( info=info, dataset_path="/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical", ) - if not new_path.exists(): - new_path.parent.mkdir(parents=True, exist_ok=True) - bf.symlink_files(new_path, hard_link=True) # exist_ok=True, - print(new_path) + # Build the set of source extensions we'd actually hard-link. + srcs: dict[str, Path] = {ext: src for ext, src in bf.file.items() if Path(src).exists()} + # For segmentations (msk), skip auto-generated grid-only JSON sidecars. + if bf.format == "msk" and "json" in srcs and _is_grid_only_json(srcs["json"]): + srcs.pop("json") + if not srcs: + continue + # Compute the real target paths (per extension) and skip if all already exist. + base = str(new_path)[: -len(".nii.gz")] + targets = {ext: Path(base + "." + ext) for ext in srcs} + if all(t.exists() for t in targets.values()): + continue + new_path.parent.mkdir(parents=True, exist_ok=True) + # Temporarily restrict bf.file to just the sources we want linked, then restore. + original_file = bf.file.copy() + bf._file = srcs # bypass the property's auto-discovery (already _checked=True) + try: + bf.symlink_files(new_path, hard_link=True) + for t in targets.values(): + print(t) + finally: + bf._file = original_file leftover = {k: v for k, v in d.items() if v is not None} assert len(leftover) == 0, leftover + marker = _hard_link_done_marker(str(subj)) + marker.parent.mkdir(parents=True, exist_ok=True) + marker.touch(exist_ok=True) + + +# --------------------------------------------------------------------------- +# nako_export.py corrected-outputs integration +# --------------------------------------------------------------------------- +# +# nako_export.py writes fetswap-corrected VIBE / MEVIBE files under +# ``dataset-nako-canonical/rawdata-corrected/`` (see the docstring at the top +# of that script). Because it runs incrementally, at any given moment only +# some subjects/chunks/sequs are corrected. We track what's replaced in a +# side-JSON so downstream consumers know which files to swap for the +# corrected ones — and, for VIBE, when to re-stitch water/fat because a +# per-chunk correction breaks the pre-existing stitched raw volume. + +_CANONICAL_ROOT = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical") +_CORRECTED_ROOT = _CANONICAL_ROOT / "rawdata-corrected" +_DEFAULT_CORRECTED_INDEX = Path(__file__).with_name("_load_nako_wh_corrected.json") + + +def build_corrected_index( + corrected_root: Path = _CORRECTED_ROOT, + out_json: Path = _DEFAULT_CORRECTED_INDEX, + verbose: bool = True, +) -> dict: + """Scan ``rawdata-corrected/`` and record which VIBE chunks / MEVIBE sequs nako_export replaced. + + Layout of the produced JSON:: + + {"": {"vibe": {"corrected_chunks": [1, 2, ...]}, "mevibe": {"corrected_sequs": ["4", "5", ...]}}} + + The JSON is rewritten from scratch on every call — nako_export runs + incrementally, so we always take disk state as the truth. + """ + import re + + index: dict = {} + if not corrected_root.exists(): + log.on_warning(f"corrected root {corrected_root} does not exist; writing empty index") + else: + for sub_dir in sorted(corrected_root.glob("*/*")): + if not (sub_dir.is_dir() and sub_dir.name.isdigit()): + continue + sub = sub_dir.name + entry: dict = {} + vibe_dir = sub_dir / "vibe" + if vibe_dir.exists(): + chunks: set[int] = set() + for p in vibe_dir.glob(f"sub-{sub}_acq-ax_chunk-*_part-water_desc-corrected_vibe.nii.gz"): + m = re.search(r"chunk-(\d+)", p.name) + if m: + chunks.add(int(m.group(1))) + if chunks: + entry["vibe"] = {"corrected_chunks": sorted(chunks)} + mevibe_dir = sub_dir / "mevibe" + if mevibe_dir.exists(): + sequs: set[str] = set() + for p in mevibe_dir.glob(f"sub-{sub}_sequ-*_acq-ax_part-water_desc-corrected_mevibe.nii.gz"): + m = re.search(r"sequ-([^_]+)", p.name) + if m: + sequs.add(m.group(1)) + if sequs: + entry["mevibe"] = {"corrected_sequs": sorted(sequs)} + if entry: + index[sub] = entry + out_json.parent.mkdir(parents=True, exist_ok=True) + tmp = out_json.with_suffix(out_json.suffix + ".tmp") + tmp.write_text(json.dumps(index, indent=2, sort_keys=True)) + tmp.replace(out_json) + if verbose: + n_vibe = sum(1 for v in index.values() if "vibe" in v) + n_mevibe = sum(1 for v in index.values() if "mevibe" in v) + log.on_log(f"corrected index: {len(index)} subjects ({n_vibe} vibe, {n_mevibe} mevibe) -> {out_json}") + return index + + +def load_corrected_index(path: Path = _DEFAULT_CORRECTED_INDEX) -> dict: + """Read the JSON produced by :func:`build_corrected_index`, or ``{}`` if missing.""" + if not Path(path).exists(): + return {} + try: + return json.loads(Path(path).read_text()) + except json.JSONDecodeError: + log.on_warning(f"corrupt corrected index at {path}; ignoring") + return {} + + +def _restitch_vibe_water_fat( + sub: str, + corrected_chunks: list[int], + dataset: Path = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako"), + scratch_root: Path = _RESAMPLE_TMP_ROOT, +) -> dict[str, Path]: + """Re-stitch VIBE water & fat for ``sub`` mixing corrected chunks with raw ones. + + Any chunk listed in ``corrected_chunks`` is pulled from + ``rawdata-corrected/…/vibe/…_desc-corrected_vibe.nii.gz``; everything else + comes from ``rawdata/…/vibe/…_vibe.nii.gz``. The stitched output goes + under ``$TPTBOX_RESAMPLED_SCRATCH/rawdata_stitched/…/vibe/…``. + + The corrected-chunk set is baked into the output filename so a later + ``nako_export`` pass that corrects more chunks produces a distinct file + (no stale cache). Returns ``{"water": Path, "fat": Path}`` (partial when + stitching fails for one part). + """ + import re + + from TPTBox.stitching import stitching as _stitching_fn + + ss = sub[:3] + corr_dir = _CORRECTED_ROOT / ss / sub / "vibe" + raw_dir = dataset / "rawdata" / ss / sub / "vibe" + if not raw_dir.exists(): + log.on_warning(f"sub-{sub}: no raw vibe dir at {raw_dir}; skipping re-stitch") + return {} + + all_chunks: set[int] = set() + for p in raw_dir.glob(f"sub-{sub}_acq-ax_chunk-*_part-water_vibe.nii.gz"): + m = re.search(r"chunk-(\d+)", p.name) + if m: + all_chunks.add(int(m.group(1))) + if not all_chunks: + log.on_warning(f"sub-{sub}: no raw vibe chunks found under {raw_dir}") + return {} + corrected_set = set(corrected_chunks) & all_chunks + chunks_sorted = sorted(all_chunks) + tag = "corrected" + "".join(f"C{c}" for c in sorted(corrected_set)) + + out_dir = scratch_root / "rawdata_stitched" / ss / sub / "vibe" + out_dir.mkdir(parents=True, exist_ok=True) + + result: dict[str, Path] = {} + for part in ("water", "fat"): + out_path = out_dir / f"sub-{sub}_sequ-stitched_acq-ax_part-{part}_desc-{tag}_vibe.nii.gz" + if out_path.exists(): + result[part] = out_path + continue + images: list[Path] = [] + for c in chunks_sorted: + if c in corrected_set: + p = corr_dir / f"sub-{sub}_acq-ax_chunk-{c}_part-{part}_desc-corrected_vibe.nii.gz" + if p.exists(): + images.append(p) + continue + p = raw_dir / f"sub-{sub}_acq-ax_chunk-{c}_part-{part}_vibe.nii.gz" + if p.exists(): + images.append(p) + if len(images) < 2: + log.on_warning(f"sub-{sub} part-{part}: only {len(images)} chunk(s) available; cannot stitch") + continue + try: + _stitching_fn( + [str(p) for p in images], + str(out_path), + is_seg=False, + bias_field=False, + verbose=False, + verbose_stitching=False, + ) + except Exception as e: # noqa: BLE001 + log.on_warning(f"sub-{sub} part-{part}: stitching failed: {type(e).__name__}: {e}") + continue + result[part] = out_path + return result + + +def _corrected_bids_file(path: Path, dataset_root: Path) -> BIDS_FILE: + return BIDS_FILE(str(path), str(dataset_root), verbose=False) + + +def _apply_corrections_to_subj_dict(sub: str, subj_dict: dict, index: dict) -> dict: + """Swap ``subj_dict`` entries for their nako_export-corrected counterparts, if any. + + Returns a small report ``{"mevibe": [sequs...], "vibe": [chunks...]}`` describing + what was actually replaced (useful for logging / hard-link verification). + """ + report: dict[str, list] = {"mevibe": [], "vibe": []} + entry = index.get(str(sub)) + if not entry: + return report + + # ---- MEVIBE: whole-sequ replacement (part-fat replaces the buggy `mevibe_part-fat` slot). ---- + if "mevibe" in entry: + for sequ in entry["mevibe"]["corrected_sequs"]: + corr_dir = _CORRECTED_ROOT / str(sub)[:3] / str(sub) / "mevibe" + pdff = corr_dir / f"sub-{sub}_sequ-{sequ}_acq-ax_part-fat-fraction_desc-corrected_mevibe.nii.gz" + if pdff.exists(): + subj_dict["mevibe_part-fat"] = _corrected_bids_file(pdff, _CANONICAL_ROOT) + report["mevibe"].append(sequ) + + # ---- VIBE: re-stitch water + fat from corrected + raw chunks. ---- + if "vibe" in entry: + stitched = _restitch_vibe_water_fat(str(sub), entry["vibe"]["corrected_chunks"]) + for part in ("water", "fat"): + if part in stitched: + subj_dict[f"vibe_part-{part}"] = _corrected_bids_file(stitched[part], _RESAMPLE_TMP_ROOT) + report["vibe"] = list(entry["vibe"]["corrected_chunks"]) + + return report + + +def verify_hardlink(sub: str, corrected_index_path: Path = _DEFAULT_CORRECTED_INDEX) -> None: + """Run the loop for a single subject, apply corrections, then trace what + :func:`hard_link` *would* do and check every source exists and every target + would land on the same filesystem as its source (so :func:`os.link` won't + hit ``EXDEV``). Prints one line per file — no writes are performed. + """ + index = load_corrected_index(corrected_index_path) + sub = str(sub) + ss = sub[:3] + hits = 0 + for d in loop_over_repaired_nako(test=True, test_key=f"/{ss}/{sub}", corrected_index=index): + if str(d.get("id", "")) != sub: + continue + hits += 1 + subj_report = _apply_corrections_to_subj_dict(sub, d, index) + print(f"[verify-hardlink] sub-{sub} corrections applied: {subj_report}") + + def _check(bf, parent: str, info: dict | None = None) -> None: + if bf is None: + return + if isinstance(bf, str): + bf = BIDS_FILE(bf, "/DATA/NAS/datasets_processed/NAKO/dataset-nako/", verbose=False) + src = bf.get_nii_file() + try: + target = bf.get_changed_path( + "nii.gz", + bf.format, + parent=parent, + info=info or {}, + dataset_path=str(_CANONICAL_ROOT), + ) + except Exception as e: # noqa: BLE001 + print(f" ERR {bf}: get_changed_path failed: {e}") + return + src_exists = src is not None and Path(src).exists() + same_fs = src is not None and Path(src).stat().st_dev == _CANONICAL_ROOT.stat().st_dev if src_exists else False + status = "ok" if src_exists and same_fs else ("cross-fs" if src_exists else "src-missing") + print(f" [{status:>10s}] {src} -> {target}") + + for _key, t2w in (d.get("t2w_chunk") or {}).items(): + if t2w: + _check(t2w[0], parent="rawdata", info={"ses": "baseline"}) + seg_keys = ( + "msk_seg-body-composition_mod-mevibe", + "vibeseg100", + "MRSegmentator", + "msk_seg-body-composition_mod-vibe", + "roi", + "vert", + "spine", + "poi", + ) + img_keys = ( + "pd", + "T2haste", + "T2w", + "eco0-opp1", + "eco1-pip1", + "eco2-opp2", + "eco3-in1", + "eco4-pop1", + "eco5-arb1", + "mevibe_part-fat", + "vibe_part-outphase", + "vibe_part-fat", + "vibe_part-water", + "vibe_part-inphase", + ) + for keys, parent in ((img_keys, "rawdata"), (seg_keys, "derivatives")): + for k in keys: + _check(d.get(k), parent=parent, info={"run": None}) + if hits == 0: + print(f"[verify-hardlink] sub-{sub} was not produced by loop_over_repaired_nako (test_key filter?)") def _grid_worker(nii_path: str) -> tuple[str, str]: @@ -785,6 +1295,9 @@ def _to_path(v): if isinstance(v, (str, Path)): s = str(v) return s if s and Path(s).exists() else None + # Skip segmentations (msk-format files) — grid prewarm is only for image volumes. + if getattr(v, "format", None) == "msk": + return None get_nii_file = getattr(v, "get_nii_file", None) if get_nii_file is None: return None @@ -792,6 +1305,7 @@ def _to_path(v): p = get_nii_file() except Exception: # noqa: BLE001 return None + # JSON-only BIDS entries (e.g. fullbody POI) have no nii.gz — nothing to warm. return str(p) if p is not None else None for key, v in subj_dict.items(): @@ -846,7 +1360,7 @@ def _drain(fs): if verbose: log.on_warning(f"grid precompute {path}: {status}") - with ProcessPoolExecutor(max_workers=num_workers) as pool: + with _non_interactive_mode(), ProcessPoolExecutor(max_workers=num_workers) as pool: pending: list = [] for subj_dict in loop_over_repaired_nako(**loop_kwargs): for p in _iter_grid_targets(subj_dict): @@ -871,22 +1385,34 @@ def _drain(fs): log = Print_Logger() - parser = argparse.ArgumentParser(description="NAKO helpers: hard-link or prewarm grid info.") + parser = argparse.ArgumentParser(description="NAKO helpers: hard-link, prewarm grid info, or track nako_export corrections.") parser.add_argument( "--precompute-grid", action="store_true", help="Prewarm bf.get_grid_info() JSON caches in parallel processes instead of hard-linking.", ) + parser.add_argument( + "--build-corrected-index", + action="store_true", + help="Scan dataset-nako-canonical/rawdata-corrected/ and (re)write the corrections JSON.", + ) + parser.add_argument( + "--verify-hardlink", + metavar="SUB", + default=None, + help="Trace hard_link()'s planned links for one subject; check src exists + same fs as target.", + ) parser.add_argument("--workers", type=int, default=None, help="Number of worker processes (default: cpu_count-1).") - # parser.add_argument("--no-test", action="store_true", help="Iterate the full dataset instead of the default test subtree.") args = parser.parse_args() - test = False - if args.precompute_grid: + test = True + + if args.build_corrected_index: + build_corrected_index() + elif args.verify_hardlink is not None: + verify_hardlink(args.verify_hardlink) + elif args.precompute_grid: precompute_grid_info_parallel(num_workers=args.workers, test=test) else: - for d in loop_over_repaired_nako(test=test): - # pass + corrected = load_corrected_index() + for d in loop_over_repaired_nako(test=test, corrected_index=corrected, skip_subject=is_hard_linked): hard_link(d) - # print(d["T2w"]) - # break - # check VIBE same shape From bd99bb65437d892eaccb7505d1400772db280b3b Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 14 Sep 2026 11:03:59 +0000 Subject: [PATCH 09/26] add prelimenary scoliose and pelic paramerters --- TPTBox/spine/spinestats/README.md | 139 ++++- TPTBox/spine/spinestats/_load_nako.py | 1 + TPTBox/spine/spinestats/_qc_report.py | 228 +++++++ TPTBox/spine/spinestats/_run_all.py | 269 ++++++++- .../spine/spinestats/all_output_reference.md | 104 +++- TPTBox/spine/spinestats/curvature.py | 556 ++++++++++++++++++ TPTBox/spine/spinestats/pelvic_parameters.py | 237 ++++++++ 7 files changed, 1516 insertions(+), 18 deletions(-) create mode 100644 TPTBox/spine/spinestats/_qc_report.py create mode 100644 TPTBox/spine/spinestats/curvature.py create mode 100644 TPTBox/spine/spinestats/pelvic_parameters.py diff --git a/TPTBox/spine/spinestats/README.md b/TPTBox/spine/spinestats/README.md index 3282494e..5afa6ae9 100644 --- a/TPTBox/spine/spinestats/README.md +++ b/TPTBox/spine/spinestats/README.md @@ -8,6 +8,8 @@ objects and `NII` segmentations. | Module | Description | |---|---| | `angles.py` | Cobb angle, cervical lordosis, thoracic kyphosis, lumbar lordosis | +| `curvature.py` | Extended curvature metrics: SVA, coronal balance, wedge angles, segmental endplate angles, axial rotation, spline-based curvature profile, multi-curve Cobb | +| `pelvic_parameters.py` | Pelvic Incidence / Pelvic Tilt / Sacral Slope / PI-LL mismatch from the fullbody-POI json | | `measure_ivd_and_vertebra_geometry.py` | Per-structure geometry (heights, widths, x1–x6) and T2 signal ratio for vertebrae and IVDs | | `torso_vat_sat.py` | VBQ score, body composition CSA, muscle fat infiltration, torso VAT/SAT/muscle volumes; also `peak_centered_mean` | | `vertebra_anatomical_widths.py` | Anatomical distances per vertebra (IVD height, body height, LR/AP widths) stored on `POI.info` | @@ -55,13 +57,20 @@ not on the Python API. A standalone copy of this reference lives at | Key | Source function | What it covers | |---|---|---| -| `ivd_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=100)` | intervertebral discs | -| `vert_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=50)` | vertebral bodies | +| `ivd_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=100)` | intervertebral discs (now also carries wedge angles/indices per label) | +| `vert_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=50)` | vertebral bodies (now also carries wedge angles/indices per label) | | `VBQ_score` | `VBQ_score` | vertebral bone quality (T2 signal ratio) | | `body_composition_score` | `body_composition_score` | axial CSA per tissue at chosen vertebral levels | | `muscle_fat_infiltration` | `muscle_fat_infiltration` | Dixon fat-fraction based muscle-quality metrics | | `torso_vat_sat_muscle_mass` | `torso_vat_sat_muscle_mass` | whole-torso VAT / SAT / muscle volume | | `cobb`, `curv` | `plot_cobb_and_lordosis_and_kyphosis` | only when called with `cobb=True` | +| `sva` | `curvature.compute_sva` | Sagittal Vertical Axis (mm) | +| `coronal_balance` | `curvature.compute_coronal_balance` | Coronal Balance (mm) | +| `axial_rotation` | `curvature.compute_axial_rotation` | per-vertebra axial rotation angle | +| `segmental_endplate_angles` | `curvature.compute_segmental_endplate_angles` | inter-vertebral wedge (disc) angle | +| `curvature_profile` | `curvature.compute_curvature_profile` | spline-based arc/chord/κ profile with apex positions | +| `multi_cobb` | `curvature.compute_multi_cobb` | multi-curve Cobb detection from coronal spline projection | +| `pelvic_parameters` | `pelvic_parameters.compute_pelvic_parameters` | PI / PT / SS + PI-LL mismatch in multiple variants | Distance metrics from `vertebra_anatomical_widths.compute_all_distances` are not currently written into the json by `run_all`; they live on the @@ -302,6 +311,132 @@ Implementation notes: Values can be `None` if the required vertebrae are missing from the POI. +## Extended curvature metrics (`curvature.py`) + +Everything below is written into the json by `run_all` when +`need_curvature=True` (default). All angles are in **degrees**, lengths +in **millimetres**. On missing landmarks the corresponding entry +contains `None` values plus an `error` message; the pipeline never +raises for these. + +### `sva` — Sagittal Vertical Axis + +Signed horizontal offset in the sagittal plane between the top vertebra +(default C7) and a base reference. Positive = top vertebra is anterior +of the base (typical adult). + +- `sva_mm` — the offset in mm +- `top_vertebra`, `base_vertebra`, `base_landmark` — which vertebrae / + landmark were used (falls back S1 → L5 → L4 if the earlier is + missing) +- `top_pi_coords`, `base_pi_coords` — (P, I) coordinates of the two + points in the internal POI orientation, for QC + +**Caveat:** measured on supine MRI. Standing SVA is typically 0-50 mm +larger; comparisons to Schwab-style thresholds derived from standing +radiographs are only approximate. + +### `coronal_balance` + +Signed horizontal offset in the coronal plane between the top vertebra +(C7) and the base vertebra R coordinate (CSVL proxy). Positive = top +vertebra is right of CSVL. + +- `coronal_balance_mm` +- `top_vertebra`, `base_vertebra` +- `top_r_coord`, `base_r_coord` + +### `axial_rotation` + +`dict[vertebra_name, degrees]` (e.g. `"L1": -3.5`). Signed angle in the +axial plane between `Vertebra_Direction_Right` and the image right +axis. Positive = rotation towards the patient's left. + +### `segmental_endplate_angles` + +`dict["-", degrees]`. Signed sagittal-plane angle between +the inferior endplate direction of the upper vertebra and the inferior +endplate direction of the lower vertebra. Positive = anterior opening +(typical lordotic disc). + +### `curvature_profile` + +Spline fit through the `Vertebra_Corpus` centroids (uses +`POI.fit_spline`, cubic B-spline). Reports: + +- `arc_length_mm`, `chord_length_mm`, `tortuosity` (arc/chord) +- `curvature_max_1_per_mm`, `curvature_mean_1_per_mm` — |κ| in 3D +- `curvature_sagittal_max_1_per_mm`, `curvature_coronal_max_1_per_mm` + — |κ| in the two 2D projections +- `apices`, `sagittal_apices`, `coronal_apices` — each a list of up to + 6 dicts `{arc_mm, kappa_1_per_mm}` ordered by arc position + +### `multi_cobb` + +Automatic multi-curve Cobb detection from the coronal spline +projection. Sign changes of the signed curvature are treated as +inflection points; between each pair of consecutive inflections one +Cobb angle is reported. + +- `curves`: list of `{arc_start_mm, arc_end_mm, apex_arc_mm, length_mm, + cobb_deg, handedness}` (handedness = `"right"` or `"left"`) +- `max_cobb_deg`: maximum |Cobb| across all detected curves + +### Wedge metrics on vert_geometry / ivd_geometry + +`compute_wedge_metrics` merges four extra fields **into each label's +entry** of `vert_geometry` and `ivd_geometry` (so they automatically +flow into `per_vertebra.xlsx` / `per_ivd.xlsx`): + +- `sagittal_wedge_deg` — `atan((x1 − x2) / x6)`, positive = anterior taller +- `coronal_wedge_deg` — `atan((x3 − x4) / x5)`, positive = right taller +- `sagittal_wedge_index` — `(x1 − x2) / mean(x1, x2)`, unitless +- `coronal_wedge_index` — `(x3 − x4) / mean(x3, x4)`, unitless + +Genant-style fracture screening: a `sagittal_wedge_index` below about +`-0.4` corresponds to > 40 % anterior height loss. + +## Pelvic parameters (`pelvic_parameters.py`) + +Written under the `pelvic_parameters` key when a fullbody-POI json is +available under +`/derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_..._seg-fullbody_poi.json`. + +Two variants are always computed side by side so they can be compared +in QC. The `poi_ap` variant is expected to be the canonical one; the +`poi_ala` variant uses a laterally-averaged reference that in most +subjects deviates enough to serve as a robustness check. + +Each variant reports: + +- `pi_deg` (Pelvic Incidence, unsigned; anatomical constant) +- `pt_deg` (Pelvic Tilt, signed; positive = sacrum posterior of hip axis) +- `ss_deg` (Sacral Slope, unsigned; endplate tilt from horizontal) +- `pi_ll_mismatch_deg` (`pi_deg − lumbar_lordosis`, using the pipeline's + supine LL). `None` if LL is missing. +- `hip_center_mm`, `s1_endplate_center_mm` — the two 3D points used, for QC + +Relationship: **PI = PT + SS** (up to sign convention). If the two +sides disagree by more than a fraction of a degree the landmarks are +inconsistent. + +**Limitations** (in `pelvic_parameters.py` module docstring): + +1. **Supine vs. standing.** Metrics are derived from supine MRI. + Standing SS is typically ~10-15° larger, standing PT ~10-15° smaller + than the same subject supine. **PI is anatomical** and comparable + across positions. PI-LL uses the supine LL and is therefore not + directly comparable to Schwab thresholds derived from standing images. +2. The S1 upper endplate is reconstructed from two point landmarks + (`Sacral_Crest_S1` posterior + `Anterior_Longitudinal_Medial` + anterior for `poi_ap`); the ligament attachment can drift inferior + with age / degeneration and bias the endplate normal. +3. The bi-femoral axis uses the atlas-registered `PELVIS_CENTER` + landmark that lives under `femur_right` / `femur_left` in the + fullbody-POI json. +4. No axial pelvic obliquity correction: the sagittal plane is world + `(y, z)`. In practice supine subjects are close to aligned. + --- ## `vertebra_anatomical_widths.py` diff --git a/TPTBox/spine/spinestats/_load_nako.py b/TPTBox/spine/spinestats/_load_nako.py index 38e1a8d3..b0462ee2 100644 --- a/TPTBox/spine/spinestats/_load_nako.py +++ b/TPTBox/spine/spinestats/_load_nako.py @@ -102,6 +102,7 @@ def get_corrected_mevibe(fam: BIDS_Family, compute_PDFF=True): # TODO return di def get_current_best_T2w_seg(sub, black_list_t2w=None): if black_list_t2w is None: + # 111007 Needs fix, very strong scolisose black_list_t2w = [ # Head missing T2w "106910", diff --git a/TPTBox/spine/spinestats/_qc_report.py b/TPTBox/spine/spinestats/_qc_report.py new file mode 100644 index 00000000..fd5eb02a --- /dev/null +++ b/TPTBox/spine/spinestats/_qc_report.py @@ -0,0 +1,228 @@ +"""Generate a QC report for the aggregated NAKO Excel outputs. + +Produces one .xlsx file next to the input tables with: +- one summary sheet listing coverage, missing/error counts, distributions, and + the count of "outlier" values per metric (values outside a hard-coded + physiological / plausible range); +- one sheet per outlier metric listing the subjects (or subject+label rows) + that fall outside that range so they can be reviewed by hand. + +Usage +----- + python -m TPTBox.spine.spinestats._qc_report [] + +If no folder is passed, defaults to +``/DATA/NAS/ongoing_projects/robert/test/NAKO-stats``. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pandas as pd + +# -------------------------------------------------------------------------- +# Outlier ranges (inclusive). A value outside this range is flagged for review. +# These are chosen wide enough to keep almost all real biological variation, +# so what remains is very likely a segmentation / landmark artifact. +# -------------------------------------------------------------------------- + +# per_subject metrics +OUTLIER_SUBJECT: dict[str, tuple[float, float]] = { + "sva.sva_mm": (-80, 100), + "coronal_balance.coronal_balance_mm": (-80, 80), + "curvature_profile.arc_length_mm": (350, 700), + "curvature_profile.tortuosity": (1.0, 1.20), + "curvature_profile.curvature_max_1_per_mm": (0.0, 0.05), + "multi_cobb.max_cobb_deg": (0, 60), + "curv.cervical_lordosis": (0, 80), + "curv.thoracic_kyphosis": (0, 80), + "curv.lumbar_lordosis": (0, 80), + "pelvic_parameters.poi_ap.pi_deg": (10, 90), + "pelvic_parameters.poi_ap.pt_deg": (-30, 60), + "pelvic_parameters.poi_ap.ss_deg": (-15, 75), + "pelvic_parameters.poi_ap.pi_ll_mismatch_deg": (-40, 40), + # NAKO T2w calibration shifts VBQ ~0.2 lower than clinical values from the + # literature. Empirical NAKO distribution: median 0.28, IQR [0.24, 0.33]. + # Range chosen wide enough to keep all normal biological variation while + # flagging clear segmentation / signal-normalization failures. + "VBQ_score.VBQ_L1-L4": (0.1, 0.8), +} + +# per_vertebra metrics (also apply to per_ivd where the column exists) +OUTLIER_LABEL: dict[str, tuple[float, float]] = { + "axial_rotation_deg": (-30, 30), + "endplate_internal_angle_deg": (0, 45), + "sagittal_wedge_deg": (-30, 30), + "coronal_wedge_deg": (-20, 20), + "sagittal_wedge_index": (-0.6, 0.6), + "coronal_wedge_index": (-0.6, 0.6), + "segmental_endplate_angle_deg": (-30, 30), +} + +TOP_OFFENDERS = 50 + + +# -------------------------------------------------------------------------- +# Report builders +# -------------------------------------------------------------------------- + + +def _summary_row(df: pd.DataFrame, col: str, lo: float, hi: float) -> dict: + v = pd.to_numeric(df[col], errors="coerce") + n_total = len(df) + n_valid = int(v.notna().sum()) + v_ok = v.dropna() + mask_out = (v_ok < lo) | (v_ok > hi) + n_out = int(mask_out.sum()) + if n_valid == 0: + return { + "column": col, "n_total": n_total, "n_valid": 0, "n_missing": n_total, + "median": None, "iqr_low": None, "iqr_high": None, "min": None, "max": None, + "outlier_range": f"[{lo}, {hi}]", "n_outliers": 0, "pct_outliers": None, + } + return { + "column": col, + "n_total": n_total, + "n_valid": n_valid, + "n_missing": n_total - n_valid, + "median": round(float(v_ok.median()), 3), + "iqr_low": round(float(v_ok.quantile(0.25)), 3), + "iqr_high": round(float(v_ok.quantile(0.75)), 3), + "min": round(float(v_ok.min()), 3), + "max": round(float(v_ok.max()), 3), + "outlier_range": f"[{lo}, {hi}]", + "n_outliers": n_out, + "pct_outliers": round(100.0 * n_out / n_valid, 3), + } + + +def _outlier_frame( + df: pd.DataFrame, col: str, lo: float, hi: float, extra_cols: list[str] +) -> pd.DataFrame: + v = pd.to_numeric(df[col], errors="coerce") + mask = v.notna() & ((v < lo) | (v > hi)) + if not mask.any(): + return pd.DataFrame(columns=["subject", col, *extra_cols]) + dfo = df.loc[mask, [c for c in ["subject", "label", col, *extra_cols] if c in df.columns]].copy() + dfo = dfo.sort_values(col, key=lambda s: pd.to_numeric(s, errors="coerce").abs(), ascending=False) + return dfo.head(TOP_OFFENDERS) + + +def build_qc_report(folder: Path) -> Path: + sub_p = folder / "per_subject.xlsx" + vert_p = folder / "per_vertebra.xlsx" + ivd_p = folder / "per_ivd.xlsx" + out_p = folder / "qc_report.xlsx" + + print(f"loading {sub_p} ...", flush=True) + sub = pd.read_excel(sub_p) + print(f"loading {vert_p} ...", flush=True) + vert = pd.read_excel(vert_p) + print(f"loading {ivd_p} ...", flush=True) + ivd = pd.read_excel(ivd_p) + + # ------------------------------------------------------------------ + # Header sheet + # ------------------------------------------------------------------ + header = pd.DataFrame( + [ + {"item": "n_subjects", "value": len(sub)}, + {"item": "per_subject_cols", "value": len(sub.columns)}, + {"item": "n_vertebra_rows", "value": len(vert)}, + {"item": "n_ivd_rows", "value": len(ivd)}, + {"item": "pelvic_error_rate_%", + "value": round(100.0 * sub.get("pelvic_parameters.error", pd.Series([np.nan] * len(sub))).notna().sum() / len(sub), 3)}, + ] + ) + + # PI = PT + SS invariant check + if "pelvic_parameters.poi_ap.pi_deg" in sub.columns: + pi = pd.to_numeric(sub["pelvic_parameters.poi_ap.pi_deg"], errors="coerce") + pt = pd.to_numeric(sub["pelvic_parameters.poi_ap.pt_deg"], errors="coerce") + ss = pd.to_numeric(sub["pelvic_parameters.poi_ap.ss_deg"], errors="coerce") + diff = (pi - (pt + ss)).abs() + n_valid = int(pi.notna().sum()) + n_ok = int((diff < 0.1).sum()) + header = pd.concat( + [ + header, + pd.DataFrame( + [ + {"item": "PI_valid", "value": n_valid}, + {"item": "PI=PT+SS_within_0.1_deg", "value": n_ok}, + {"item": "PI=PT+SS_violation_pct", "value": round(100.0 * (n_valid - n_ok) / max(n_valid, 1), 3)}, + ] + ), + ], + ignore_index=True, + ) + + # ------------------------------------------------------------------ + # Per-column summary + # ------------------------------------------------------------------ + sub_summary = pd.DataFrame( + [_summary_row(sub, c, lo, hi) for c, (lo, hi) in OUTLIER_SUBJECT.items() if c in sub.columns] + ) + vert_summary = pd.DataFrame( + [_summary_row(vert, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in vert.columns] + ) + ivd_summary = pd.DataFrame( + [_summary_row(ivd, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in ivd.columns] + ) + + # ------------------------------------------------------------------ + # Outlier rows + # ------------------------------------------------------------------ + print("writing report ...", flush=True) + with pd.ExcelWriter(out_p, engine="xlsxwriter") as w: + header.to_excel(w, sheet_name="_header", index=False) + sub_summary.to_excel(w, sheet_name="summary_subject", index=False) + vert_summary.to_excel(w, sheet_name="summary_vertebra", index=False) + ivd_summary.to_excel(w, sheet_name="summary_ivd", index=False) + + # Subject outliers + for col, (lo, hi) in OUTLIER_SUBJECT.items(): + if col not in sub.columns: + continue + dfo = _outlier_frame(sub, col, lo, hi, extra_cols=[]) + if not dfo.empty: + sn = _sheet_name(col, prefix="s_") + dfo.to_excel(w, sheet_name=sn, index=False) + + # Per-vertebra outliers + for col, (lo, hi) in OUTLIER_LABEL.items(): + if col not in vert.columns: + continue + dfo = _outlier_frame(vert, col, lo, hi, extra_cols=["label"]) + if not dfo.empty: + sn = _sheet_name(col, prefix="v_") + dfo.to_excel(w, sheet_name=sn, index=False) + + # Per-IVD outliers + for col, (lo, hi) in OUTLIER_LABEL.items(): + if col not in ivd.columns: + continue + dfo = _outlier_frame(ivd, col, lo, hi, extra_cols=["label"]) + if not dfo.empty: + sn = _sheet_name(col, prefix="i_") + dfo.to_excel(w, sheet_name=sn, index=False) + + print(f"wrote {out_p} ({out_p.stat().st_size} bytes)") + return out_p + + +def _sheet_name(col: str, prefix: str = "") -> str: + r"""Excel sheet names must be <=31 chars and cannot contain []:*?/\ .""" + s = prefix + col.replace(".", "_").replace(":", "_") + for ch in "[]:*?/\\": + s = s.replace(ch, "_") + return s[:31] + + +if __name__ == "__main__": + default = Path("/DATA/NAS/ongoing_projects/robert/test/NAKO-stats") + folder = Path(sys.argv[1]) if len(sys.argv) > 1 else default + build_qc_report(folder) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index d26f1bcf..236d189b 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -42,6 +42,14 @@ "body_composition_score", "muscle_fat_infiltration", "torso_vat_sat_muscle_mass", + # Extended curvature / balance / pelvic metrics (see curvature.py and pelvic_parameters.py). + "sva", + "coronal_balance", + "axial_rotation", + "segmental_endplate_angles", + "curvature_profile", + "multi_cobb", + "pelvic_parameters", ) @@ -85,6 +93,9 @@ def get_nako_paths(nako_id: str) -> dict[str, Path | None]: vert, spine = v, sp break + fullbody_poi = ( + DATASET_ROOT / f"derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_sequ-stitched_acq-ax_part-water_seg-fullbody_poi.json" + ) roi = ( DATASET_ROOT / f"derivatives_Abdominal-Segmentation/{pfx}/{sub}/vibe/sub-{nako_id}_sequ-stitched_acq-ax_mod-vibe_seg-ROI_msk.nii.gz" ) @@ -100,6 +111,7 @@ def get_nako_paths(nako_id: str) -> dict[str, Path | None]: "spine": spine, "roi": roi, "vibeseg100": vibeseg100 if vibeseg100.exists() else None, + "fullbody_poi": fullbody_poi if fullbody_poi.exists() else None, "dataset": DATASET_ROOT, } @@ -152,6 +164,8 @@ def run_all( need_bcs=True, need_mfi=True, need_torso=True, + need_curvature=True, + need_pelvic=True, ) -> dict[str, Any] | None: """Run the full pipeline for one subject and return the results dict. @@ -184,9 +198,19 @@ def run_all( """ from TPTBox import Location, calc_poi_from_subreg_vert from TPTBox.spine.spinestats.angles import plot_cobb_and_lordosis_and_kyphosis + from TPTBox.spine.spinestats.curvature import ( + compute_axial_rotation, + compute_coronal_balance, + compute_curvature_profile, + compute_multi_cobb, + compute_segmental_endplate_angles, + compute_sva, + compute_wedge_metrics, + ) from TPTBox.spine.spinestats.measure_ivd_and_vertebra_geometry import ( measure_ivd_and_vertebra_geometry, # structure_label: int = 100 and structure_label: int = 49 ) + from TPTBox.spine.spinestats.pelvic_parameters import compute_pelvic_parameters from TPTBox.spine.spinestats.torso_vat_sat import VBQ_score, body_composition_score, muscle_fat_infiltration, torso_vat_sat_muscle_mass if "t2w" not in file_dict: @@ -234,6 +258,9 @@ def _need(*keys: str, compute: bool) -> bool: need_bcs = _need("body_composition_score", compute=need_bcs) need_mfi = _need("muscle_fat_infiltration", compute=need_mfi) need_torso = _need("torso_vat_sat_muscle_mass", compute=need_torso) + _curvature_keys = ("sva", "coronal_balance", "axial_rotation", "segmental_endplate_angles", "curvature_profile", "multi_cobb") + need_curvature = _need(*_curvature_keys, compute=need_curvature) + need_pelvic = _need("pelvic_parameters", compute=need_pelvic) # Recompute area save = False if "VBQ_score" in out and "VBQ_L1-L1_old" in out["VBQ_score"]: @@ -246,8 +273,42 @@ def _need(*keys: str, compute: bool) -> bool: need_torso = True del out["torso_vat_sat_muscle_mass"] save = True + # Redo pelvic_parameters if the cached entry is just an error stub (e.g. old runs where + # fullbody_poi wasn't resolved), OR if it was written with the old unsigned-SS + # implementation (detected via the PI = PT + SS invariant violated by > 0.1 deg). + if "pelvic_parameters" in out: + pp = out.get("pelvic_parameters", {}) + redo = False + if isinstance(pp, dict): + if pp.get("fullbody_poi_json") is None and "error" in pp: + redo = True + else: + for variant_key in ("poi_ap", "poi_ala"): + v = pp.get(variant_key) + if not isinstance(v, dict): + continue + pi = v.get("pi_deg") + pt = v.get("pt_deg") + ss = v.get("ss_deg") + if pi is None or pt is None or ss is None: + continue + try: + if abs(float(pi) - (float(pt) + float(ss))) > 0.1: + redo = True + break + except (TypeError, ValueError): + continue + if redo: + need_pelvic = True + del out["pelvic_parameters"] + save = True + # Retrospectively fold per-vertebra metrics into vert_geometry / ivd_geometry + # for cached JSONs written before the merge was in place. Cheap and idempotent. + if any(k in out for k in ("axial_rotation", "endplate_internal_angle", "segmental_endplate_angles")): + _merge_per_vertebra_metrics(out) + save = True #### - need_poi = need_cobb or need_ivd or need_vert + need_poi = need_cobb or need_ivd or need_vert or need_curvature need_t2w = need_ivd or need_vert or need_vbq need_vert_nii = need_poi or need_vbq or need_bcs or need_mfi need_spine_nii = need_vert_nii or need_vbq @@ -255,7 +316,7 @@ def _need(*keys: str, compute: bool) -> bool: need_roi = need_mfi or need_torso need_vibe_wf = need_mfi - if not (need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso): + if not (need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso or need_curvature or need_pelvic): if _merge_endplate_angles(out, Path(poi_out)) or save: save_json(final_out, out) return out @@ -329,6 +390,66 @@ def _need(*keys: str, compute: bool) -> bool: logger.on_debug("torso_vat_sat_muscle_mass") torso_results, _body_comp = torso_vat_sat_muscle_mass(vibe_seg, roi, dataset_id=100) out["torso_vat_sat_muscle_mass"] = torso_results + + if need_curvature and poi is not None: + try: + logger.on_debug("curvature metrics") + out["sva"] = compute_sva(poi) + out["coronal_balance"] = compute_coronal_balance(poi) + out["axial_rotation"] = compute_axial_rotation(poi) + out["segmental_endplate_angles"] = compute_segmental_endplate_angles(poi) + out["curvature_profile"] = compute_curvature_profile(poi) + out["multi_cobb"] = compute_multi_cobb(poi) + except Exception: + logger.on_fail("curvature error caught") + logger.print_error() + + # Merge wedge metrics directly into the per-label vert_geometry / ivd_geometry + # entries so they land in per_vertebra.xlsx / per_ivd.xlsx automatically. + for geom_key in ("vert_geometry", "ivd_geometry"): + geom = out.get(geom_key) + if not isinstance(geom, dict): + continue + try: + geom_int_keys = {int(k): v for k, v in geom.items()} + wedge = compute_wedge_metrics(geom_int_keys) + for label, w in wedge.items(): + target = geom.get(str(label)) or geom.get(label) + if isinstance(target, dict): + for k, v in w.items(): + target.setdefault(k, v) + except Exception: + logger.on_fail(f"wedge merge failed for {geom_key}") + logger.print_error() + + # Also fold per-vertebra dicts (axial_rotation, endplate_internal_angle) into + # vert_geometry entries and per-IVD segmental angles into ivd_geometry, so they + # land in per_vertebra.xlsx / per_ivd.xlsx rather than exploding per_subject + # into dozens of extra columns. + _merge_per_vertebra_metrics(out) + + if need_pelvic: + try: + from TPTBox.spine.spinestats.pelvic_parameters import resolve_fullbody_poi_path + + lumbar_ll = None + curv = out.get("curv") + if isinstance(curv, dict): + lumbar_ll = curv.get("lumbar_lordosis") + # Prefer an explicit path from file_dict (get_nako_paths sets one); + # fall back to resolving from (dataset, id) so subjects streamed from + # loop_over_repaired_nako (which doesn't add fullbody_poi) still work. + fb = file_dict.get("fullbody_poi") + if fb is None: + sub_id = file_dict.get("id") + ds = file_dict.get("dataset", DATASET_ROOT) + if sub_id is not None: + fb = resolve_fullbody_poi_path(ds, str(sub_id)) + out["pelvic_parameters"] = compute_pelvic_parameters(fb, lumbar_lordosis_deg=lumbar_ll) + except Exception: + logger.on_fail("pelvic_parameters error caught") + logger.print_error() + _merge_endplate_angles(out, Path(poi_out)) logger.on_save("save", final_out.name) save_json(final_out, out) @@ -357,6 +478,59 @@ def _read_endplate_internal_angles(poi_json_path: Path) -> dict[str, Any]: return {} +def _merge_per_vertebra_metrics(out: dict[str, Any]) -> None: + """Fold per-vertebra top-level dicts into vert_geometry / ivd_geometry entries. + + Moves values from: + - ``axial_rotation`` {vertebra_name: deg} -> vert_geometry[label]["axial_rotation_deg"] + - ``endplate_internal_angle`` {vertebra_name: deg} -> vert_geometry[label]["endplate_internal_angle_deg"] + - ``segmental_endplate_angles`` {"V1-V2": deg} -> ivd_geometry[100+V1_label]["segmental_endplate_angle_deg"] + + Non-destructive on the top-level dicts (kept for backward reads), but + the collector will exclude these keys from per_subject.xlsx. + """ + from TPTBox.core.vert_constants import Vertebra_Instance + + def _name_to_label(n: str) -> int | None: + try: + return Vertebra_Instance[n].value + except KeyError: + return None + + def _find(geom: dict, label: int) -> dict | None: + return geom.get(str(label)) or geom.get(label) + + vg = out.get("vert_geometry") + if isinstance(vg, dict): + for src_key, dst_key in ( + ("axial_rotation", "axial_rotation_deg"), + ("endplate_internal_angle", "endplate_internal_angle_deg"), + ): + src = out.get(src_key) + if not isinstance(src, dict): + continue + for name, val in src.items(): + lab = _name_to_label(str(name)) + if lab is None: + continue + target = _find(vg, lab) + if isinstance(target, dict): + target.setdefault(dst_key, val) + + ig = out.get("ivd_geometry") + if isinstance(ig, dict): + seg = out.get("segmental_endplate_angles") + if isinstance(seg, dict): + for pair, val in seg.items(): + upper = str(pair).split("-", 1)[0] + lab = _name_to_label(upper) + if lab is None: + continue + target = _find(ig, 100 + lab) + if isinstance(target, dict): + target.setdefault("segmental_endplate_angle_deg", val) + + def _merge_endplate_angles(out: dict[str, Any], poi_json_path: Path) -> bool: """Attach the POI's per-vertebra endplate_internal_angle to ``out``. @@ -422,7 +596,16 @@ def _rows_from_json(subject_id: str, data: dict) -> tuple[dict[str, Any], list[d row limit (1_048_576). """ per_subject: dict[str, Any] = {"subject": subject_id} - subject_view = {k: v for k, v in data.items() if k not in ("ivd_geometry", "vert_geometry")} + # Exclude per-label geometry dicts (their rows live in per_vertebra / per_ivd) + # and per-vertebra dicts that were already merged into vert_geometry / ivd_geometry. + _PER_SUBJECT_EXCLUDE = ( + "ivd_geometry", + "vert_geometry", + "axial_rotation", + "endplate_internal_angle", + "segmental_endplate_angles", + ) + subject_view = {k: v for k, v in data.items() if k not in _PER_SUBJECT_EXCLUDE} _flatten("", subject_view, per_subject) per_vert: list[dict[str, Any]] = [] @@ -457,21 +640,62 @@ def _collector_worker( ivd_rows: list[dict[str, Any]] = [] seen: set[str] = set() - def _flush() -> None: - if subject_rows: - pd.DataFrame(subject_rows).to_excel(out_folder / per_subject_name, index=False) - if vertebra_rows: - pd.DataFrame(vertebra_rows).to_excel(out_folder / per_vertebra_name, index=False) - if ivd_rows: - pd.DataFrame(ivd_rows).to_excel(out_folder / per_ivd_name, index=False) + log_path = out_folder / "excel_collector.log" + + def _log(msg: str) -> None: + try: + with log_path.open("a") as f: + from datetime import datetime as _dt + + f.write(f"[{_dt.now().isoformat(timespec='seconds')}] {msg}\n") + except Exception: + pass + + def _write(df_rows: list[dict[str, Any]], name: str) -> None: + if not df_rows: + return + target = out_folder / name + tmp = target.with_suffix(target.suffix + ".tmp") + # xlsxwriter is ~5-10x faster than openpyxl for wide sheets; fall back if unavailable. + engine: str | None = "xlsxwriter" + try: + import xlsxwriter # noqa: F401 + except ImportError: + engine = None + _log(f"writing {name}: rows={len(df_rows)} engine={engine or 'openpyxl'}") + try: + pd.DataFrame(df_rows).to_excel(tmp, index=False, engine=engine) + tmp.replace(target) + _log(f" {name} done: {target.stat().st_size} bytes") + except Exception as e: + _log(f" {name} FAILED: {type(e).__name__}: {e}") + try: + tmp.unlink(missing_ok=True) + except Exception: + pass + + def _flush(final: bool) -> None: + # All three tables are now written only at shutdown. The mid-run + # per_subject flush was killing the daemon on large runs (silent + # xlsxwriter/openpyxl crash during full-DataFrame rewrites), so + # nothing is written mid-run — the log below still emits a heartbeat + # every ``flush_every`` subjects so progress stays visible. + if not final: + return + _write(subject_rows, per_subject_name) + _write(vertebra_rows, per_vertebra_name) + _write(ivd_rows, per_ivd_name) + _log(f"collector started; flush_every={flush_every} (heartbeat only; all writes at shutdown)") while True: try: item = task_q.get(timeout=1.0) except _queue.Empty: continue if item is None: - _flush() + _log(f"final flush triggered; seen={len(seen)} vertebra_rows={len(vertebra_rows)} ivd_rows={len(ivd_rows)}") + _flush(final=True) + _log("collector exiting") return subject_id, json_path = item if subject_id in seen: @@ -486,7 +710,7 @@ def _flush() -> None: ivd_rows.extend(per_ivd) seen.add(subject_id) if flush_every and len(seen) % flush_every == 0: - _flush() + _log(f"heartbeat: seen={len(seen)} vertebra_rows={len(vertebra_rows)} ivd_rows={len(ivd_rows)}") class ExcelCollector: @@ -541,7 +765,15 @@ def submit(self, subject_id: str, json_path: str | Path) -> None: raise RuntimeError("ExcelCollector not started") self._queue.put((str(subject_id), str(json_path))) - def close(self, join_timeout: float = 60.0) -> None: + def close(self, join_timeout: float = 1800.0) -> None: + """Signal the daemon to flush + exit and wait up to ``join_timeout`` seconds. + + The final flush writes ``per_vertebra.xlsx`` and ``per_ivd.xlsx`` + from scratch; with 30k subjects that can take 5-15 minutes per + file. The default timeout is generous (30 min) so the daemon + has enough time to finish the shutdown flush. Progress is + logged to ``/excel_collector.log``. + """ if self._proc is None: return self._queue.put(None) @@ -617,6 +849,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: aggregate = True do_not_update = False test = False + collector: ExcelCollector | None = None if aggregate: collector = ExcelCollector(out_folder=OUT_FOLDER) collector.start() @@ -669,7 +902,15 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: collector.submit(sub_id, _final_json_path(f)) finally: - if aggregate: + if aggregate and collector is not None: collector.close() if missing_rows: pd.DataFrame(missing_rows).to_excel(OUT_FOLDER / "missing_inputs.xlsx", index=False) + # Auto-generate the QC report next to the aggregated tables. + try: + from TPTBox.spine.spinestats._qc_report import build_qc_report + + build_qc_report(OUT_FOLDER) + except Exception as e: # noqa: BLE001 + logger.on_fail(f"qc_report generation failed: {type(e).__name__}: {e}") + logger.print_error() diff --git a/TPTBox/spine/spinestats/all_output_reference.md b/TPTBox/spine/spinestats/all_output_reference.md index 355a2e32..2cc16a09 100644 --- a/TPTBox/spine/spinestats/all_output_reference.md +++ b/TPTBox/spine/spinestats/all_output_reference.md @@ -12,13 +12,19 @@ Python API. | Key | Source function | What it covers | |---|---|---| -| `ivd_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=100)` | intervertebral discs | -| `vert_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=50)` | vertebral bodies | +| `ivd_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=100)` | intervertebral discs (per-label entries carry wedge angles/indices) | +| `vert_geometry` | `measure_ivd_and_vertebra_geometry(..., structure_label=50)` | vertebral bodies (per-label entries carry wedge angles/indices) | | `VBQ_score` | `VBQ_score` | vertebral bone quality (T2 signal ratio) | | `body_composition_score` | `body_composition_score` | axial CSA per tissue at chosen vertebral levels | | `muscle_fat_infiltration` | `muscle_fat_infiltration` | Dixon fat-fraction based muscle-quality metrics | | `torso_vat_sat_muscle_mass` | `torso_vat_sat_muscle_mass` | whole-torso VAT / SAT / muscle volume | | `cobb`, `curv` | `plot_cobb_and_lordosis_and_kyphosis` | only when called with `cobb=True` | +| `sva`, `coronal_balance` | `curvature.compute_sva`, `compute_coronal_balance` | plumb-line balance offsets in mm (sagittal/coronal) | +| `axial_rotation` | `curvature.compute_axial_rotation` | per-vertebra axial rotation angle in the axial plane | +| `segmental_endplate_angles` | `curvature.compute_segmental_endplate_angles` | inter-vertebral (disc) wedge angle in the sagittal plane | +| `curvature_profile` | `curvature.compute_curvature_profile` | spline arc/chord/κ profile plus apex positions | +| `multi_cobb` | `curvature.compute_multi_cobb` | multi-curve Cobb from the coronal spline projection | +| `pelvic_parameters` | `pelvic_parameters.compute_pelvic_parameters` | PI / PT / SS + PI-LL mismatch, two variants | Angles are in **degrees**, lengths in **millimetres**, areas in **mm²**, volumes in **mm³**, fat fractions are **unitless** in `[0, 1]`, MR signal @@ -244,6 +250,100 @@ Implementation notes: Values can be `None` if the required vertebrae are missing from the POI. +## `sva` (Sagittal Vertical Axis) + +Signed horizontal offset (mm) in the sagittal plane between the top +vertebra (default C7) and a base reference (S1 → L5 → L4 fallback). +Positive = C7 anterior of the base (typical). Additional fields +`top_vertebra`, `base_vertebra`, `base_landmark`, `top_pi_coords`, +`base_pi_coords`. **Caveat:** supine MRI values differ from standing +X-ray by ~10 mm. + +## `coronal_balance` + +Signed horizontal offset (mm) in the coronal plane between C7 and the +CSVL (approximated by the base vertebra R-coordinate). Positive = C7 +right of CSVL. Extra fields: `top_vertebra`, `base_vertebra`, +`top_r_coord`, `base_r_coord`. + +## `axial_rotation` + +`dict[vertebra_name, degrees]`. Signed angle between +`Vertebra_Direction_Right` and the image right axis in the axial plane. +Positive = rotation towards the patient's left. + +## `segmental_endplate_angles` + +`dict["-", degrees]`. Signed sagittal-plane wedge angle +between the inferior endplates of two adjacent vertebrae. Approximates +the disc wedge without needing the disc mesh. + +## `curvature_profile` + +Cubic B-spline through the `Vertebra_Corpus` centroids. Reports arc +length, chord length, tortuosity (arc/chord), and |κ| statistics in +3D as well as the two 2D projections. `apices` (3D), `sagittal_apices`, +`coronal_apices` are lists of up to 6 dicts `{arc_mm, kappa_1_per_mm}` +ordered by arc position (superior → inferior). + +## `multi_cobb` + +Multi-curve Cobb detection from the coronal spline projection. Sign +changes of the signed curvature act as inflection points; one Cobb +angle is emitted per segment. + +- `curves`: `[{arc_start_mm, arc_end_mm, apex_arc_mm, length_mm, + cobb_deg, handedness}]` +- `max_cobb_deg`: max absolute Cobb across curves + +Handedness is `"right"` or `"left"` referring to the direction of the +curve's concavity. + +## Wedge fields on `vert_geometry` / `ivd_geometry` + +Added per label: + +- `sagittal_wedge_deg` — `atan((x1 − x2) / x6)` (positive = anterior taller) +- `coronal_wedge_deg` — `atan((x3 − x4) / x5)` (positive = right taller) +- `sagittal_wedge_index` — `(x1 − x2) / mean(x1, x2)` +- `coronal_wedge_index` — `(x3 − x4) / mean(x3, x4)` + +Genant-style fracture screening: `sagittal_wedge_index` below ≈ −0.4 +corresponds to > 40 % anterior height loss. + +## `pelvic_parameters` + +Present when the fullbody-POI json for the subject exists under +`derivatives-fullbody-poi/…/vibe/sub-*_seg-fullbody_poi.json`. Two +variants side by side (compare in QC): + +- `poi_ap` — canonical: uses `Sacral_Crest_S1` posterior + `Anterior_Longitudinal_Medial` anterior for the S1 endplate +- `poi_ala` — alternate: uses the midpoint of `Sacrum_Ala_Superior_L/R` as the "anterior" reference. Included as a robustness check; in practice it under-estimates PI compared to `poi_ap` + +Per variant: + +- `pi_deg` (unsigned, anatomical constant) +- `pt_deg` (signed, positive = sacrum posterior of hip axis) +- `ss_deg` (unsigned, endplate tilt from horizontal) +- `pi_ll_mismatch_deg` = `pi_deg − curv["lumbar_lordosis"]`; `None` + if lumbar lordosis is missing +- `hip_center_mm`, `s1_endplate_center_mm` for QC + +Relationship: `PI = PT + SS` (up to sign). + +**Limitations:** + +1. **Supine vs. standing.** SS/PT are position-dependent — supine SS + is systematically lower than standing SS by ~10-15°, PT + correspondingly higher. **PI is position-invariant** and the safe + number to compare across cohorts. PI-LL uses the supine LL and is + not directly comparable to Schwab-style standing thresholds. +2. S1 endplate is reconstructed from two point landmarks; the anterior + ligament attachment can drift inferior with age / degeneration. +3. Bi-femoral axis uses the atlas-registered `PELVIS_CENTER` landmark + under each femur (label 13/113 in the fullbody-POI mapping). +4. No axial-pelvic-obliquity correction; sagittal plane = world `(y, z)`. + --- ## Excel collector diff --git a/TPTBox/spine/spinestats/curvature.py b/TPTBox/spine/spinestats/curvature.py new file mode 100644 index 00000000..8f8f267f --- /dev/null +++ b/TPTBox/spine/spinestats/curvature.py @@ -0,0 +1,556 @@ +"""Spine curvature metrics beyond Cobb and regional lordosis/kyphosis. + +All functions operate on a :class:`TPTBox.POI` in the internal +``("P", "I", "R")`` orientation (posterior/inferior/right axes 0/1/2), +which the pipeline already produces via ``poi.reorient_().rescale_()``. +Coordinates therefore live in millimetres. + +Contents +-------- +- :func:`compute_sva` + Sagittal Vertical Axis (mm): horizontal distance in the sagittal + plane between the C7 plumb line and the posterior-superior corner + of S1. Positive = C7 is anterior of S1 (typical adult spine). +- :func:`compute_coronal_balance` + Coronal Balance (mm): horizontal distance between the C7 plumb line + and the Central Sacral Vertical Line (midpoint of the sacrum). + Positive = C7 is to the patient's right of the CSVL. +- :func:`compute_wedge_metrics` + Rewrites the raw x1..x6 heights already stored in + ``vert_geometry``/``ivd_geometry`` into anterior/posterior and + left/right wedge angles (in degrees) and wedge indices (unitless). +- :func:`compute_segmental_endplate_angles` + Inter-vertebral (segmental) wedge angles between the inferior + endplate of vertebra N and the superior endplate of vertebra N+1. + Approximates disc wedging without needing the disc mesh. +- :func:`compute_axial_rotation` + Per-vertebra axial rotation angle (degrees) between + ``Vertebra_Direction_Right`` and the image-space right axis, in the + axial (P-R) plane. Positive = rotation towards the patient's left. +- :func:`compute_curvature_profile` + B-spline through the ``Vertebra_Corpus`` centroids. Returns arc + length, chord length, tortuosity, sampled |κ|(s), and the arc + positions plus |κ| values of the largest curvature peaks (apices). +- :func:`compute_multi_cobb` + Multi-curve scoliotic Cobb detection based on the curvature profile: + finds inflection points in the coronal projection and returns one + Cobb angle per curve between neighbouring inflections. +- :data:`EXTENDED_CURVATURE_DEFINITIONS` + Additional entries for ``angles.curvature_definition``: + ``t1_slope``, ``c2_c7_angle``, ``cervical_sva_helper``. Merge into + the default map when you want them in the lordosis/kyphosis output. + +Design notes +------------ +- Every function is total: on missing landmarks it returns a dict with + the metric keys set to ``None``/``NaN`` and an ``"error"`` message. +- No side effects on the input POI beyond ``reorient_().rescale_()`` + (which the pipeline already does). +- The Cobb/lordosis code in ``angles.py`` stays untouched. +""" + +from __future__ import annotations + +from itertools import pairwise +from typing import Any + +import numpy as np +from numpy.linalg import norm + +from TPTBox import POI, Location, Vertebra_Instance +from TPTBox.spine.spinestats.angles import Def_Curvature, MoveTo + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +# Internal POI axes after reorient/rescale: 0=P, 1=I, 2=R. +_AX_P, _AX_I, _AX_R = 0, 1, 2 + + +def _prep(poi: POI) -> POI: + """Return the POI in the internal ``(P, I, R)`` orientation at 1 mm scale.""" + return poi.reorient().rescale() + + +def _get(poi: POI, vert: int | Vertebra_Instance, loc: Location) -> np.ndarray | None: + """Fetch a coordinate as ``np.array`` or ``None`` if not present.""" + v = vert.value if isinstance(vert, Vertebra_Instance) else vert + key = (v, loc.value if isinstance(loc, Location) else loc) + if key not in poi: + return None + return np.asarray(poi[key], dtype=float) + + +def _corpus(poi: POI, vert: int | Vertebra_Instance) -> np.ndarray | None: + return _get(poi, vert, Location.Vertebra_Corpus) + + +def _endplate(poi: POI, vert: int | Vertebra_Instance, superior: bool) -> np.ndarray | None: + """Return the (approximate) endplate center of ``vert``. + + Uses the standard ``Endplate`` subregion. If it's not there for both + endplates individually, fall back to the corpus + inferior direction + to estimate the endplate midpoint. + """ + loc = Location.Vertebral_Body_Endplate_Superior if superior else Location.Vertebral_Body_Endplate_Inferior + p = _get(poi, vert, loc) + if p is not None: + return p + corp = _corpus(poi, vert) + inf = _get(poi, vert, Location.Vertebra_Direction_Inferior) + if corp is None or inf is None: + return None + d = inf - corp + n = norm(d) + if n == 0: + return corp + return corp + (d / n) * (5.0 if not superior else -5.0) + + +def _last_present(poi: POI, candidates: list[Vertebra_Instance]) -> Vertebra_Instance | None: + for v in candidates: + if _corpus(poi, v) is not None: + return v + return None + + +# --------------------------------------------------------------------------- +# 1. Sagittal Vertical Axis (SVA) +# --------------------------------------------------------------------------- + + +_SVA_BASE_FALLBACKS = [Vertebra_Instance.S1, Vertebra_Instance.L5, Vertebra_Instance.L4] + + +def compute_sva(poi: POI, top_vert: Vertebra_Instance = Vertebra_Instance.C7) -> dict[str, Any]: + """Sagittal Vertical Axis: signed horizontal offset (mm) in the sagittal plane. + + Definition + ---------- + Drop a plumb line from the center of the ``top_vert`` (default C7) + vertebra body downwards (image-inferior axis). Measure the signed + distance along the posterior-anterior axis to the posterior-superior + corner of S1 (falls back to L5, then L4, if S1 is absent). Positive + values → top vertebra is anterior of the base reference (positive + sagittal balance, typical adult). + + Returns: + ------- + dict + - ``sva_mm``: signed offset in mm (or ``None`` if landmarks missing) + - ``top_vertebra`` / ``base_vertebra`` / ``base_landmark`` + - ``top_pi_coords`` / ``base_pi_coords`` + - ``error``: message if computation failed + """ + poi = _prep(poi) + top = _corpus(poi, top_vert) + if top is None: + return {"sva_mm": None, "error": f"{top_vert.name} corpus missing"} + + base_vert = None + base_ref = None + landmark = None + for cand in _SVA_BASE_FALLBACKS: + corp = _corpus(poi, cand) + if corp is None: + continue + ep = _endplate(poi, cand, superior=True) + base_vert = cand + base_ref = ep if ep is not None else corp + landmark = f"{cand.name}_endplate_superior" if ep is not None else f"{cand.name}_corpus" + break + if base_ref is None or base_vert is None: + return {"sva_mm": None, "error": "no base vertebra (S1/L5/L4) available"} + + sva = float(base_ref[_AX_P] - top[_AX_P]) + return { + "sva_mm": round(sva, 2), + "top_vertebra": top_vert.name, + "base_vertebra": base_vert.name, + "base_landmark": landmark, + "top_pi_coords": [round(float(top[_AX_P]), 2), round(float(top[_AX_I]), 2)], + "base_pi_coords": [round(float(base_ref[_AX_P]), 2), round(float(base_ref[_AX_I]), 2)], + } + + +# --------------------------------------------------------------------------- +# 2. Coronal Balance +# --------------------------------------------------------------------------- + + +def compute_coronal_balance(poi: POI, top_vert: Vertebra_Instance = Vertebra_Instance.C7) -> dict[str, Any]: + """Coronal Balance: signed horizontal offset (mm) in the coronal plane. + + Definition + ---------- + Distance along the patient's right axis between ``top_vert`` (default + C7) and the Central Sacral Vertical Line, taken here as the + R-coordinate of the S1 corpus (falls back to L5, then L4). + Positive → top vertebra is to the patient's right of the CSVL. + + Returns: + ------- + dict + - ``coronal_balance_mm`` + - ``top_vertebra`` / ``base_vertebra`` + - ``top_r_coord`` / ``base_r_coord`` + - ``error`` + """ + poi = _prep(poi) + top = _corpus(poi, top_vert) + if top is None: + return {"coronal_balance_mm": None, "error": f"{top_vert.name} corpus missing"} + base = None + base_vert = None + for cand in _SVA_BASE_FALLBACKS: + c = _corpus(poi, cand) + if c is not None: + base = c + base_vert = cand + break + if base is None or base_vert is None: + return {"coronal_balance_mm": None, "error": "no base vertebra (S1/L5/L4) available"} + cb = float(top[_AX_R] - base[_AX_R]) + return { + "coronal_balance_mm": round(cb, 2), + "top_vertebra": top_vert.name, + "base_vertebra": base_vert.name, + "top_r_coord": round(float(top[_AX_R]), 2), + "base_r_coord": round(float(base[_AX_R]), 2), + } + + +# --------------------------------------------------------------------------- +# 3. + 5. Wedge metrics on top of the existing x1..x6 measurements +# --------------------------------------------------------------------------- + + +def compute_wedge_metrics(geometry: dict[int, dict[str, float]]) -> dict[int, dict[str, float]]: + """Wedge angles and indices derived from the already-computed x1..x6. + + Given a ``vert_geometry`` or ``ivd_geometry`` section (label → + metrics with ``anterior_height_x1``, ``posterior_height_x2``, + ``right_height_x3``, ``left_height_x4``, ``width_sagittal_x6``, + ``width_lateral_x5``), this returns per label: + + - ``sagittal_wedge_deg`` = atan((x1 − x2) / x6) (positive = anterior taller) + - ``coronal_wedge_deg`` = atan((x3 − x4) / x5) (positive = right taller) + - ``sagittal_wedge_index`` = (x1 − x2) / mean(x1, x2) + - ``coronal_wedge_index`` = (x3 − x4) / mean(x3, x4) + + Labels with missing / non-positive inputs get NaN for that metric. + Non-destructive: returns a fresh dict. + """ + out: dict[int, dict[str, float]] = {} + for label, m in geometry.items(): + if not isinstance(m, dict): + continue + x1 = m.get("anterior_height_x1") + x2 = m.get("posterior_height_x2") + x3 = m.get("right_height_x3") + x4 = m.get("left_height_x4") + x5 = m.get("width_lateral_x5") + x6 = m.get("width_sagittal_x6") + + def _finite(v): + try: + return v is not None and np.isfinite(float(v)) + except (TypeError, ValueError): + return False + + row: dict[str, float] = {} + if _finite(x1) and _finite(x2) and _finite(x6) and float(x6) > 0: + row["sagittal_wedge_deg"] = round(float(np.degrees(np.arctan2(float(x1) - float(x2), float(x6)))), 3) + avg = (float(x1) + float(x2)) / 2.0 + row["sagittal_wedge_index"] = round((float(x1) - float(x2)) / avg, 4) if avg > 0 else np.nan + else: + row["sagittal_wedge_deg"] = np.nan + row["sagittal_wedge_index"] = np.nan + + if _finite(x3) and _finite(x4) and _finite(x5) and float(x5) > 0: + row["coronal_wedge_deg"] = round(float(np.degrees(np.arctan2(float(x3) - float(x4), float(x5)))), 3) + avg = (float(x3) + float(x4)) / 2.0 + row["coronal_wedge_index"] = round((float(x3) - float(x4)) / avg, 4) if avg > 0 else np.nan + else: + row["coronal_wedge_deg"] = np.nan + row["coronal_wedge_index"] = np.nan + + out[label] = row + return out + + +# --------------------------------------------------------------------------- +# 4. Segmental endplate angles (inter-vertebral wedge) +# --------------------------------------------------------------------------- + + +def compute_segmental_endplate_angles(poi: POI) -> dict[str, float]: + """Angle (degrees) between adjacent vertebral endplates in the sagittal plane. + + For each pair of neighbouring vertebrae with both endplate normals + available, computes the signed angle between (a) the inferior + endplate line of vertebra N and (b) the superior endplate line of + vertebra N+1, projected onto the sagittal plane. Approximates disc + wedging without needing the disc mesh. + + Endplate line: perpendicular to ``Vertebra_Direction_Inferior`` + (which is normal to the endplate), in the sagittal plane. + + Returns: + ------- + dict keyed by ``"-"``, e.g. ``"L4-L5"``. + Positive value → anterior opening (typical lordotic disc). + """ + poi = _prep(poi) + order = Vertebra_Instance.order_dict() + ordered = sorted( + [v for v in Vertebra_Instance if _get(poi, v, Location.Vertebra_Direction_Inferior) is not None], + key=lambda v: order.get(v.value, v.value), + ) + out: dict[str, float] = {} + for a, b in pairwise(ordered): + if b.value - a.value not in (1,) and not (a.value < 25 and b.value < 25 and b.value == a.value + 1): + # Only measure real neighbours (skip jumps like L5→S1 numbering gap). + pass + ca = _corpus(poi, a) + cb = _corpus(poi, b) + ia = _get(poi, a, Location.Vertebra_Direction_Inferior) + ib = _get(poi, b, Location.Vertebra_Direction_Inferior) + if ca is None or cb is None or ia is None or ib is None: + continue + # Endplate normal (inferior direction) → project onto sagittal (P-I) plane. + na = np.array([ia[_AX_P] - ca[_AX_P], ia[_AX_I] - ca[_AX_I]]) + nb = np.array([ib[_AX_P] - cb[_AX_P], ib[_AX_I] - cb[_AX_I]]) + na_n = norm(na) + nb_n = norm(nb) + if na_n == 0 or nb_n == 0: + continue + na /= na_n + nb /= nb_n + # Signed angle between the two inferior normals in the sagittal plane. + cross = na[0] * nb[1] - na[1] * nb[0] + dot = float(np.clip(na @ nb, -1.0, 1.0)) + ang = float(np.degrees(np.arctan2(cross, dot))) + out[f"{a.name}-{b.name}"] = round(ang, 3) + return out + + +# --------------------------------------------------------------------------- +# 6. Axial rotation per vertebra +# --------------------------------------------------------------------------- + + +def compute_axial_rotation(poi: POI) -> dict[str, float]: + """Per-vertebra axial rotation in the axial (P-R) plane, in degrees. + + Definition: signed angle between ``Vertebra_Direction_Right`` (from + the vertebral body's centroid) and the image-space right axis, + measured in the axial plane. 0° = right-direction aligned with image + R axis, positive = rotation toward patient's left (counterclockwise + when viewed from superior). + + Requires ``Vertebra_Direction_Right`` and ``Vertebra_Corpus`` in the POI. + """ + poi = _prep(poi) + out: dict[str, float] = {} + for v in Vertebra_Instance: + c = _corpus(poi, v) + r = _get(poi, v, Location.Vertebra_Direction_Right) + if c is None or r is None: + continue + d = r - c + # Axial plane = (P, R) plane; drop the I component. + planar = np.array([d[_AX_P], d[_AX_R]]) + n = norm(planar) + if n == 0: + continue + planar /= n + # Reference axis = image right = (0, 1) in (P, R). + ang = float(np.degrees(np.arctan2(planar[0], planar[1]))) + out[v.name] = round(ang, 3) + return out + + +# --------------------------------------------------------------------------- +# 7. Curvature profile via spline + apex detection +# --------------------------------------------------------------------------- + + +def _numerical_curvature(points: np.ndarray) -> np.ndarray: + """|κ(s)| along an equidistantly-sampled curve (2D or 3D), via finite differences.""" + if len(points) < 3: + return np.zeros(len(points)) + dp = np.gradient(points, axis=0) + ddp = np.gradient(dp, axis=0) + # 2D branch: signed curvature magnitude. 3D branch: cross-product norm. + num = np.abs(dp[:, 0] * ddp[:, 1] - dp[:, 1] * ddp[:, 0]) if points.shape[1] == 2 else norm(np.cross(dp, ddp), axis=1) + den = norm(dp, axis=1) ** 3 + with np.errstate(divide="ignore", invalid="ignore"): + k = np.where(den > 0, num / den, 0.0) + return k + + +def compute_curvature_profile( + poi: POI, + smoothness: int = 10, + samples_per_poi: int = 20, + top_k_apices: int = 6, +) -> dict[str, Any]: + """Curvature profile of the spine, based on the internal ``fit_spline``. + + Fits a cubic B-spline through the ``Vertebra_Corpus`` centroids + (sorted by ``Vertebra_Instance.order_dict()``) and returns: + + - ``arc_length_mm`` / ``chord_length_mm`` / ``tortuosity`` (arc/chord) + - ``curvature_max_1_per_mm`` / ``curvature_mean_1_per_mm`` + - ``apices``: list of ``{arc_mm, kappa_1_per_mm}`` for the ``top_k_apices`` + local maxima of |κ|, sorted by arc position (superior → inferior). + - ``sagittal_apices`` / ``coronal_apices``: same, but on the + projected 2D curve in each plane (better matches the clinical + notion of "the apex of a scoliotic curve"). + + All units follow the POI: mm. + """ + poi = _prep(poi) + # fit_spline expects at least a handful of points. + if len(poi.extract_subregion(Location.Vertebra_Corpus)) < 4: + return {"error": "too few Vertebra_Corpus POIs to fit a spline"} + try: + pts, _der = poi.fit_spline( + smoothness=smoothness, + samples_per_poi=samples_per_poi, + location=Location.Vertebra_Corpus, + vertebra=True, + ) + except Exception as e: + return {"error": f"fit_spline failed: {e}"} + + seg_len = norm(np.diff(pts, axis=0), axis=1) + arc = np.concatenate([[0.0], np.cumsum(seg_len)]) + arc_length = float(arc[-1]) + chord_length = float(norm(pts[-1] - pts[0])) + tortuosity = arc_length / chord_length if chord_length > 0 else np.nan + + kappa_3d = _numerical_curvature(pts) + kappa_sag = _numerical_curvature(pts[:, [_AX_P, _AX_I]]) + kappa_cor = _numerical_curvature(pts[:, [_AX_R, _AX_I]]) + + apices_3d = _find_apices(arc, kappa_3d, top_k_apices) + apices_sag = _find_apices(arc, kappa_sag, top_k_apices) + apices_cor = _find_apices(arc, kappa_cor, top_k_apices) + + return { + "arc_length_mm": round(arc_length, 2), + "chord_length_mm": round(chord_length, 2), + "tortuosity": round(float(tortuosity), 5) if np.isfinite(tortuosity) else None, + "curvature_max_1_per_mm": round(float(kappa_3d.max()), 6), + "curvature_mean_1_per_mm": round(float(kappa_3d.mean()), 6), + "curvature_sagittal_max_1_per_mm": round(float(kappa_sag.max()), 6), + "curvature_coronal_max_1_per_mm": round(float(kappa_cor.max()), 6), + "apices": apices_3d, + "sagittal_apices": apices_sag, + "coronal_apices": apices_cor, + } + + +def _find_apices(arc: np.ndarray, kappa: np.ndarray, top_k: int) -> list[dict[str, float]]: + """Return up to ``top_k`` local maxima of ``kappa`` as ``{arc_mm, kappa_...}`` dicts.""" + if len(kappa) < 3: + return [] + peak_mask = (kappa[1:-1] > kappa[:-2]) & (kappa[1:-1] > kappa[2:]) + peak_idx = np.flatnonzero(peak_mask) + 1 + if len(peak_idx) == 0: + return [] + peak_idx = peak_idx[np.argsort(-kappa[peak_idx])[:top_k]] + peak_idx = np.sort(peak_idx) + return [{"arc_mm": round(float(arc[i]), 2), "kappa_1_per_mm": round(float(kappa[i]), 6)} for i in peak_idx] + + +# --------------------------------------------------------------------------- +# 8. Multi-curve Cobb via inflection points +# --------------------------------------------------------------------------- + + +def compute_multi_cobb(poi: POI, min_curve_length_mm: float = 30.0) -> dict[str, Any]: + """Automatic multi-curve Cobb detection from the coronal spline projection. + + Fits the ``Vertebra_Corpus`` spline (as in :func:`compute_curvature_profile`) + and projects it onto the coronal (R, I) plane. Finds inflection + points as sign changes of the coronal signed curvature; between each + pair of consecutive inflections, the local Cobb angle is measured as + the angle between the tangent at the start and end of that segment. + + Curves shorter than ``min_curve_length_mm`` are dropped as noise. + + Returns: + ------- + dict + - ``curves``: list of dicts with ``arc_start_mm``, ``arc_end_mm``, + ``apex_arc_mm``, ``length_mm``, ``cobb_deg``, and + ``handedness`` (``"right"``/``"left"`` = direction of the + curve's concavity). + - ``max_cobb_deg``: maximum |Cobb| across all detected curves. + """ + poi = _prep(poi) + if len(poi.extract_subregion(Location.Vertebra_Corpus)) < 4: + return {"error": "too few Vertebra_Corpus POIs"} + try: + pts, _der = poi.fit_spline(location=Location.Vertebra_Corpus, vertebra=True, smoothness=10, samples_per_poi=20) + except Exception as e: + return {"error": f"fit_spline failed: {e}"} + + # Project onto coronal plane (R, I): axis 2 = R (x), axis 1 = I (y). + cor = pts[:, [_AX_R, _AX_I]] + dp = np.gradient(cor, axis=0) + ddp = np.gradient(dp, axis=0) + signed_k = dp[:, 0] * ddp[:, 1] - dp[:, 1] * ddp[:, 0] + seg_len = norm(np.diff(pts, axis=0), axis=1) + arc = np.concatenate([[0.0], np.cumsum(seg_len)]) + + sign = np.sign(signed_k) + change = np.flatnonzero(sign[1:] != sign[:-1]) + 1 + boundaries = np.concatenate([[0], change, [len(pts) - 1]]) + + curves: list[dict[str, Any]] = [] + for a, b in pairwise(boundaries): + if arc[b] - arc[a] < min_curve_length_mm: + continue + # Tangents at each end. + t0 = dp[a] + t1 = dp[b] + n0, n1 = norm(t0), norm(t1) + if n0 == 0 or n1 == 0: + continue + cos = float(np.clip((t0 @ t1) / (n0 * n1), -1.0, 1.0)) + cobb = float(np.degrees(np.arccos(cos))) + # Apex = |signed_k| max within the segment. + seg_k = np.abs(signed_k[a : b + 1]) + apex_local = int(np.argmax(seg_k)) + a + handedness = "right" if signed_k[apex_local] > 0 else "left" + curves.append( + { + "arc_start_mm": round(float(arc[a]), 2), + "arc_end_mm": round(float(arc[b]), 2), + "apex_arc_mm": round(float(arc[apex_local]), 2), + "length_mm": round(float(arc[b] - arc[a]), 2), + "cobb_deg": round(cobb, 3), + "handedness": handedness, + } + ) + max_cobb = max((c["cobb_deg"] for c in curves), default=0.0) + return {"curves": curves, "max_cobb_deg": round(max_cobb, 3)} + + +# --------------------------------------------------------------------------- +# 8b. Additional curvature definitions (cervical / T1-slope) +# --------------------------------------------------------------------------- + +EXTENDED_CURVATURE_DEFINITIONS: dict[str, Def_Curvature] = { + # T1 slope: angle of the T1 superior endplate to horizontal — approximated + # as the "lordosis" measured between T1 top and T1 bottom. + "t1_slope": Def_Curvature(Vertebra_Instance.T1, MoveTo.TOP, Vertebra_Instance.T1, MoveTo.BOTTOM), + # C2-C7 angle: the classic cervical Cobb between C2 inferior endplate + # and C7 inferior endplate. + "c2_c7_angle": Def_Curvature(Vertebra_Instance.C2, MoveTo.BOTTOM, Vertebra_Instance.C7, MoveTo.BOTTOM), +} diff --git a/TPTBox/spine/spinestats/pelvic_parameters.py b/TPTBox/spine/spinestats/pelvic_parameters.py new file mode 100644 index 00000000..9f449ec5 --- /dev/null +++ b/TPTBox/spine/spinestats/pelvic_parameters.py @@ -0,0 +1,237 @@ +"""Pelvic sagittal parameters (PI / PT / SS / PI-LL) from fullbody POIs. + +Reads the per-subject fullbody POI json produced by the ``TReg`` fullbody +registration pipeline +(``.../derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_..._seg-fullbody_poi.json``) +and returns the three classical pelvic parameters plus their pairwise +mismatches with lumbar lordosis. + +Multiple variants are computed side by side (see ``variants`` below). The +idea is not to *pick* one here but to expose all reasonable definitions +so they can be compared against each other (and against literature) in +downstream QC. + +Definitions +----------- +- **Pelvic Incidence (PI)**: angle between (a) the line from the + bi-coxo-femoral axis (midpoint of the two femoral head centers) to + the center of the S1 superior endplate, and (b) the perpendicular to + the S1 superior endplate. Measured in the sagittal plane. Position- + invariant (anatomical constant per subject). +- **Sacral Slope (SS)**: angle between the S1 superior endplate line + and horizontal, in the sagittal plane. Position-dependent. +- **Pelvic Tilt (PT)**: angle between the line from the bi-coxo-femoral + axis to the center of the S1 superior endplate and the vertical. + Signed positive when the sacrum is *posterior* to the hip axis + (retroverted pelvis). Position-dependent. +- Fundamental relationship: **PI = PT + SS** (up to sign convention). +- **PI-LL mismatch**: PI - lumbar_lordosis. A value close to zero is + associated with balanced spinopelvic alignment; large positive + mismatches with sagittal decompensation. + +Coordinate system +----------------- +The fullbody POI json is written in ``nib`` (nibabel world) coordinates, +so ``[x, y, z] = [right, anterior, superior]`` in mm. All computations +here therefore project onto the ``(y, z)`` sagittal plane. + +Limitations +----------- +1. **Supine vs. standing.** These metrics are conventionally measured on + standing lateral radiographs. All NAKO acquisitions are **supine** + MRI. Under gravity the sacrum tilts anteriorly; in supine SS is + systematically ~10-15° lower and PT ~10-15° higher than in the same + subject standing. **PI is anatomical and position-invariant**, so + only PI (and PI-LL when using a supine LL) are directly comparable + to standing-image reference ranges. +2. **Endplate proxy.** The S1 superior endplate is not landmarked as a + contour — it is reconstructed from two point landmarks per variant. + The ``poi_ap`` variant uses ``Sacral_Crest_S1`` (posterior) and + ``Anterior_Longitudinal_Medial`` (anterior). The anterior ligament + attachment point can drift inferior with age / degeneration, biasing + the endplate normal. +3. **Bi-femoral axis.** Uses the atlas-registered ``PELVIS_CENTER`` + landmark (which despite the name lives inside each femur landmark + group and marks the femoral head center). The registration is + template-based; large hip pathology can distort this point. +4. **PI-LL uses the supine LL** produced by the existing spine + pipeline. This is by definition smaller than a standing LL and the + mismatch numbers cannot be interpreted the same way as Schwab-style + thresholds derived from standing radiographs. +5. **No axial pelvic obliquity correction.** The sagittal plane is + taken as the world ``(y, z)`` plane. If the subject is rotated in + the scanner (obliquity around the SI axis), the projection is off + by that angle. In practice supine MRI subjects are close to aligned. + +The function is total: on missing landmarks / json each variant's block +contains ``None``s plus an ``error`` message; the outer function never +raises. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import numpy as np +from numpy.linalg import norm + +# Fullbody POI json coordinate convention: nibabel world = (R, A, S). +_IDX_R, _IDX_A, _IDX_S = 0, 1, 2 + + +def _sag(v: np.ndarray) -> np.ndarray: + """Project a nibabel-world (R, A, S) point onto the sagittal (anterior, superior) plane.""" + return v[[_IDX_A, _IDX_S]] + + +def _pelvic_from_endplate_pair(S_post: np.ndarray, S_ant: np.ndarray, FH_R: np.ndarray, FH_L: np.ndarray) -> dict[str, float]: + """Compute PI, PT, SS from posterior/anterior S1 endplate points and the two femoral head centers. + + Sign convention (Legaye / Schwab): + - SS positive when the S1 endplate tilts down anteriorly (typical). + - PT positive when the sacrum center is posterior of the hip axis + (retroverted pelvis). PT is negative for anteverted pelves. + - PI = PT + SS exactly (both in signed degrees). + + All three angles are computed as signed values via ``atan2`` in the + sagittal (anterior=Y, superior=Z) plane; ``abs()`` is avoided so the + fundamental invariant holds. + """ + S = 0.5 * (S_post + S_ant) + F = 0.5 * (FH_R + FH_L) + S2, F2 = _sag(S), _sag(F) + e = _sag(S_ant) - _sag(S_post) # posterior -> anterior in (Y, Z) + if norm(e) == 0: + return {"pi_deg": None, "pt_deg": None, "ss_deg": None, "error": "degenerate S1 endplate direction"} + d = S2 - F2 # F -> S in (Y, Z) + if norm(d) == 0: + return {"pi_deg": None, "pt_deg": None, "ss_deg": None, "error": "S1 center coincides with hip axis"} + # SS (signed): angle by which endplate tips down anteriorly. + # e = (e_y, e_z). If anterior end is inferior (e_z < 0), SS > 0. + SS = float(np.degrees(np.arctan2(-e[1], e[0]))) + # PT (signed): angle by which the F->S line tips posterior from vertical. + # d = (d_y, d_z). If S is posterior of F (d_y < 0), PT > 0 (retroverted). + PT = float(np.degrees(np.arctan2(-d[0], d[1]))) + PI = SS + PT + return { + "pi_deg": round(PI, 3), + "pt_deg": round(PT, 3), + "ss_deg": round(SS, 3), + "hip_center_mm": [round(float(x), 2) for x in F], + "s1_endplate_center_mm": [round(float(x), 2) for x in S], + } + + +def compute_pelvic_parameters( + fullbody_poi_json: Path | str | None, + lumbar_lordosis_deg: float | None = None, +) -> dict[str, Any]: + """Compute PI/PT/SS (and PI-LL) in multiple variants. + + Parameters + ---------- + fullbody_poi_json : Path + Path to ``sub-*_seg-fullbody_poi.json``. Two top-level list + entries expected (``meta``, ``body``); ``body`` maps bone → sub-index → ``[x, y, z]``. + lumbar_lordosis_deg : float, optional + The subject's lumbar lordosis (from the existing spine pipeline, + typically ``out["curv"]["lumbar_lordosis"]``). Used to compute + the ``pi_ll_mismatch_deg`` fields. + + Returns: + ------- + dict + Keys: + - ``variant`` = ``"poi_ap"``: canonical variant. Uses + ``Sacral_Crest_S1`` (posterior of S1 top) and + ``Anterior_Longitudinal_Medial`` (anterior of S1 top). + - ``variant`` = ``"poi_ala"``: alternative using midpoint of + ``Sacrum_Ala_Superior_Left/Right`` as the "anterior" reference + instead of the ligament point. Included for QC comparison + only; the ligament version is closer to the canonical + endplate midline in most subjects. + - ``pi_ll_mismatch_deg_poi_ap`` / ``..._poi_ala``: PI minus + ``lumbar_lordosis_deg`` per variant (``None`` if LL missing). + - ``fullbody_poi_json``: the resolved path used. + - ``error``: top-level message when nothing could be computed. + """ + out: dict[str, Any] = {"fullbody_poi_json": str(fullbody_poi_json) if fullbody_poi_json else None} + if fullbody_poi_json is None: + out["error"] = "no fullbody POI json path provided" + return out + p = Path(fullbody_poi_json) + if not p.exists(): + out["error"] = f"fullbody POI json does not exist: {p}" + return out + import json + + try: + with p.open() as f: + payload = json.load(f) + except Exception as e: + out["error"] = f"failed to load fullbody POI json: {e}" + return out + if not isinstance(payload, list) or len(payload) < 2 or not isinstance(payload[1], dict): + out["error"] = "unexpected fullbody POI json shape" + return out + body = payload[1] + + def _pt(bone: str, idx: str) -> np.ndarray | None: + try: + return np.asarray(body[bone][idx], dtype=float) + except (KeyError, TypeError, ValueError): + return None + + S1_post = _pt("sacrum", "1") # Sacral_Crest_S1 + S1_ant = _pt("sacrum", "19") # Anterior_Longitudinal_Medial + ala_L = _pt("sacrum", "27") # Sacrum_Ala_Superior_Left + ala_R = _pt("sacrum", "28") # Sacrum_Ala_Superior_Right + FH_R = _pt("femur_right", "11") # PELVIS_CENTER (right femoral head center) + FH_L = _pt("femur_left", "11") # PELVIS_CENTER (left femoral head center) + + missing = [] + for k, v in ( + ("sacrum_1", S1_post), + ("sacrum_19", S1_ant), + ("femur_right_11", FH_R), + ("femur_left_11", FH_L), + ): + if v is None: + missing.append(k) + if missing: + out["error"] = f"missing required landmarks: {missing}" + return out + + # Variant poi_ap. + out["poi_ap"] = _pelvic_from_endplate_pair(S1_post, S1_ant, FH_R, FH_L) # type: ignore[arg-type] + + # Variant poi_ala: use midpoint of ala_L/ala_R as the "anterior" reference. + if ala_L is not None and ala_R is not None: + ala_mid = 0.5 * (ala_L + ala_R) + out["poi_ala"] = _pelvic_from_endplate_pair(S1_post, ala_mid, FH_R, FH_L) # type: ignore[arg-type] + else: + out["poi_ala"] = {"pi_deg": None, "error": "ala landmarks missing"} + + # PI-LL mismatch per variant (when LL provided). + for variant in ("poi_ap", "poi_ala"): + v = out.get(variant, {}) + pi = v.get("pi_deg") if isinstance(v, dict) else None + if pi is not None and lumbar_lordosis_deg is not None: + v["pi_ll_mismatch_deg"] = round(float(pi) - float(lumbar_lordosis_deg), 3) + elif isinstance(v, dict): + v["pi_ll_mismatch_deg"] = None + return out + + +def resolve_fullbody_poi_path(dataset_root: Path | str, nako_id: str) -> Path | None: + """Return the expected fullbody POI json path for one NAKO subject, or ``None`` if it doesn't exist. + + Layout: + ``/derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_sequ-stitched_acq-ax_part-water_seg-fullbody_poi.json`` + where ``pfx = sub[:3]`` and ``sub = nako_id.split("_")[0].removeprefix("sub-")``. + """ + sub = str(nako_id).split("_")[0].replace("sub-", "") + pfx = sub[:3] + p = Path(dataset_root) / f"derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_sequ-stitched_acq-ax_part-water_seg-fullbody_poi.json" + return p if p.exists() else None From ba8269e94e74274bb270856f1f51cd1f19fbd50f Mon Sep 17 00:00:00 2001 From: robert Date: Mon, 14 Sep 2026 13:06:25 +0200 Subject: [PATCH 10/26] add NAKO Head scanns --- TPTBox/core/dicom/dicom_extract.py | 73 ++++++++++++++++------- TPTBox/core/dicom/dicom_header_to_keys.py | 37 ++++++++++-- 2 files changed, 82 insertions(+), 28 deletions(-) diff --git a/TPTBox/core/dicom/dicom_extract.py b/TPTBox/core/dicom/dicom_extract.py index 505d29cc..e673e01e 100644 --- a/TPTBox/core/dicom/dicom_extract.py +++ b/TPTBox/core/dicom/dicom_extract.py @@ -723,6 +723,32 @@ def _folder_fingerprint(folder: Path) -> tuple[int, str] | None: return None +def _zip_fingerprint(zip_path: Path) -> tuple[int, str] | None: + """Cheap zip identity: (1, sha1(size + mtime_ns)). + + We deliberately do NOT open the archive — for typical NAKO zips that + would take seconds per file. Size + mtime is enough to detect a fresh + re-download or a re-packed archive. + """ + import hashlib + + try: + st = zip_path.stat() + except OSError: + return None + payload = f"{st.st_size}:{st.st_mtime_ns}".encode() + return 1, hashlib.sha1(payload).hexdigest() # noqa: S324 + + +def _source_fingerprint(source: Path) -> tuple[int, str] | None: + """Dispatch to the folder- or zip-fingerprint based on the source kind.""" + if source.is_file() and source.suffix.lower() == ".zip": + return _zip_fingerprint(source) + if source.is_dir(): + return _folder_fingerprint(source) + return None + + def _extract_marker_path(source_folder: Path, dataset_path_out: Path) -> Path: """Path of the fast-skip marker for a source folder, kept under the OUTPUT dataset. @@ -735,36 +761,36 @@ def _extract_marker_path(source_folder: Path, dataset_path_out: Path) -> Path: return dataset_path_out / _EXTRACT_CACHE_DIR / f"{key}.json" -def _is_already_extracted(source_folder: Path, dataset_path_out: Path) -> bool: - """True when the source folder was extracted before and its file list is unchanged.""" +def _is_already_extracted(source: Path, dataset_path_out: Path) -> bool: + """True when the source (folder or zip) was extracted before and its fingerprint is unchanged.""" import json as _json - marker = _extract_marker_path(source_folder, dataset_path_out) + marker = _extract_marker_path(source, dataset_path_out) if not marker.is_file(): return False try: prev = _json.loads(marker.read_text()) except (OSError, ValueError): return False - fp = _folder_fingerprint(source_folder) + fp = _source_fingerprint(source) if fp is None: return False count, h = fp return prev.get("count") == count and prev.get("hash") == h -def _write_extract_marker(source_folder: Path, dataset_path_out: Path) -> None: - """Record the current file-list fingerprint for the source folder.""" +def _write_extract_marker(source: Path, dataset_path_out: Path) -> None: + """Record the current fingerprint for the source folder or zip.""" import json as _json - fp = _folder_fingerprint(source_folder) + fp = _source_fingerprint(source) if fp is None: return count, h = fp - marker = _extract_marker_path(source_folder, dataset_path_out) + marker = _extract_marker_path(source, dataset_path_out) try: marker.parent.mkdir(parents=True, exist_ok=True) - marker.write_text(_json.dumps({"source": str(source_folder), "count": count, "hash": h})) + marker.write_text(_json.dumps({"source": str(source), "count": count, "hash": h})) except OSError: pass # writing the marker is best-effort; missing it just disables the fast-path @@ -1038,18 +1064,20 @@ def extract_dicom_folder( if str(dicom_path).endswith(".pkl"): continue - # Fast-skip: identical file listing since last successful extraction - # → no DICOM headers read for this folder. Zips are excluded because - # their inner file list isn't visible without unpacking. + # Fast-skip: identical fingerprint since the last successful extraction + # → no DICOM headers read for this source. Folders fingerprint their + # rglob'd file list; zips fingerprint (size, mtime_ns) of the archive + # itself (see `_source_fingerprint`). if ( skip_already_extracted and not force_rescan - and not str(dicom_path).endswith(".zip") - and Path(dicom_path).is_dir() and _is_already_extracted(Path(dicom_path), Path(dataset_path_out)) ): logger.print(f"Skip {dicom_path} (already extracted; fingerprint matches)", verbose=verbose) continue + # Track the original source path so the marker below is keyed to the + # zip itself, not the ephemeral unpack directory. + source_for_marker = Path(dicom_path) temp_dir = None try: if str(dicom_path).endswith(".zip"): @@ -1098,14 +1126,15 @@ def process_series(key, files, parts): except Exception: logger.print_error() - # Record the fingerprint only when the whole folder went through - # without an exception AND the source is a real directory (not a - # zip mount that's about to disappear). Errors above are caught - # per-series so this fires even if individual series were skipped - # (e.g. localizers) — but not if `_read_dicom_files` itself raised - # (that path lands in the outer `finally` without reaching here). - if skip_already_extracted and temp_dir is None and Path(dicom_path).is_dir(): - _write_extract_marker(Path(dicom_path), Path(dataset_path_out)) + # Record the fingerprint only when the whole source went through + # without an exception. For zips the marker is keyed to the zip + # file itself (size+mtime) so the ephemeral temp_dir is fine; + # for folders it's keyed to the folder. Per-series errors above + # are caught inside the loop so this fires even if individual + # series were skipped (e.g. localizers) — but not if + # `_read_dicom_files` itself raised. + if skip_already_extracted: + _write_extract_marker(source_for_marker, Path(dataset_path_out)) finally: if temp_dir is not None: diff --git a/TPTBox/core/dicom/dicom_header_to_keys.py b/TPTBox/core/dicom/dicom_header_to_keys.py index 54db4da4..8d8e86f2 100644 --- a/TPTBox/core/dicom/dicom_header_to_keys.py +++ b/TPTBox/core/dicom/dicom_header_to_keys.py @@ -216,18 +216,35 @@ def _get(key, default=None): #### NAKO FIXED #### if "StudyDescription" in simp_json and "nako" in _get("StudyDescription", "").lower(): - keys["sub"] = _get("PatientID", "unnamed").split("_")[0] - series_description = _get("SeriesDescription", "unnamed") + # Read PatientID directly from simp_json — `_get` rewrites `_` to `-`, + # which would destroy the `_` split we need below. + pid_raw = str(simp_json.get("PatientID", "unnamed")).strip() + sub_part, _sep, ses_part = pid_raw.partition("_") + keys["sub"] = re.sub(r'[<>:"/\\|?*\x00-\x1F\s]', "", sub_part) or "unnamed" + # NAKO encodes the exam wave as a suffix on PatientID: + # `_30` = U1 Baseline, `_60` = U2 Follow-up. The main NAKO + # baseline export has no suffix; leave `ses` untouched there so the + # existing `use_session` (StudyDate) fallback in _get_paths still wins. + if session and ses_part: + _nako_ses_map = {"30": "baseline", "60": "followup"} + ses_clean = re.sub(r'[<>:"/\\|?*\x00-\x1F\s]', "", ses_part) + keys["ses"] = _nako_ses_map.get(ses_clean, ses_clean) + # Raw values for pattern matching — `_get` rewrites `_`→`-`, which + # would break every `T2_TSE` / `3D_GRE_TRA` / `T1_3D_SAG` check below + # and the `ProtocolName.split("_")` chunk derivation. + series_description = str(simp_json.get("SeriesDescription", "unnamed")) + protocol_name = str(simp_json.get("ProtocolName", "unnamed")) + sequ = simp_json.get("SeriesNumber") """Determine the MRI format based on the series description.""" if "T2_TSE" in series_description: - return "T2w", {"acq": "sag", "chunk": series_description.split("_")[-1], "sequ": simp_json["SeriesNumber"], **keys}, ".nii.gz" + return "T2w", {"acq": "sag", "chunk": series_description.rsplit("_", maxsplit=1)[-1], "sequ": sequ, **keys}, ".nii.gz" elif "3D_GRE_TRA" in series_description: return ( "vibe", { "acq": "ax", - "part": dixon_mapping[series_description.split("_")[-1].lower()], - "chunk": _get("ProtocolName", "unnamed").split("_")[-1], + "part": dixon_mapping[series_description.rsplit("_", maxsplit=1)[-1].lower()], + "chunk": protocol_name.rsplit("_", maxsplit=1)[-1], **keys, }, ".nii.gz", @@ -235,9 +252,17 @@ def _get(key, default=None): elif "ME_vibe" in series_description: return ( "mevibe", - {"acq": "ax", "part": dixon_mapping[series_description.split("_")[-1].lower()], "sequ": simp_json["SeriesNumber"], **keys}, + {"acq": "ax", "part": dixon_mapping[series_description.rsplit("_", maxsplit=1)[-1].lower()], "sequ": sequ, **keys}, ".nii.gz", ) + elif "T1_3D_SAG" in series_description: + # NAKO-1157 head T1 — plain sagittal ND and the MPR-Tra reformat. + acq = "tra" if "MPR_Tra" in series_description else "sag" + return "T1w", {"acq": acq, "sequ": sequ, **keys}, ".nii.gz" + elif "FLAIR" in series_description: + # NAKO-1157 head FLAIR (2D transverse). + acq = "tra" if "TRA" in series_description else "sag" + return "FLAIR", {"acq": acq, "sequ": sequ, **keys}, ".nii.gz" elif "PD" in series_description: return "pd", {"acq": "iso", **keys}, ".nii.gz" elif "T2_HASTE" in series_description: From 94bb5d6e0a98e9b67ef798d0bf9a89336e9e28d8 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Thu, 17 Sep 2026 11:51:19 +0000 Subject: [PATCH 11/26] more auto accepts --- TPTBox/spine/spinestats/_load_nako_wh.py | 203 +++++++++++++++++++---- 1 file changed, 173 insertions(+), 30 deletions(-) diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index c076ffdb..2d09d7ea 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -453,44 +453,69 @@ def get_corrected_mevibe(fam: BIDS_Family, compute_PDFF=True): # TODO return di # PDFF is recomputed out = {key: _check(fam[f"mevibe_part-{key}"]) for key in ["eco0-opp1", "eco1-pip1", "eco2-opp2", "eco3-in1", "eco4-pop1", "eco5-arb1"]} - pdff = _check(fam["mevibe_part-fat-fraction"]) if "mevibe_part-water_desc-reconstructed" in fam: - # if "mevibe_part-fat-fraction_desc-reconstructed" not in fam: fat = _check(fam["mevibe_part-fat_desc-reconstructed"]) water = _check(fam["mevibe_part-water_desc-reconstructed"]) - - else: + elif "mevibe_part-water" in fam: fat = _check(fam["mevibe_part-fat"]) water = _check(fam["mevibe_part-water"]) - out["mevibe_part-fat"] = fat - out["mevibe_part-fat"] = water - pdff = water.get_changed_bids( - "nii.gz", bids_format=water.bids_format, parent=water.parent, info={"part": "fat-fraction", "desc": "reconstructed"} + else: + fat = None + water = None + + # Anchor for BIDS derivations: prefer a raw/reconstructed water image, fall back to + # whatever fraction file the family exposes. + if water is not None: + anchor = water + elif "mevibe_part-water-fraction" in fam: + anchor = _check(fam["mevibe_part-water-fraction"]) + elif "mevibe_part-fat-fraction" in fam: + anchor = _check(fam["mevibe_part-fat-fraction"]) + else: + return out + + pdff = anchor.get_changed_bids( + "nii.gz", bids_format=anchor.bids_format, parent=anchor.parent, info={"part": "fat-fraction", "desc": "reconstructed"} ) - pdwf = water.get_changed_bids( - "nii.gz", bids_format=water.bids_format, parent=water.parent, info={"part": "water-fraction", "desc": "reconstructed"} + pdwf = anchor.get_changed_bids( + "nii.gz", bids_format=anchor.bids_format, parent=anchor.parent, info={"part": "water-fraction", "desc": "reconstructed"} ) if compute_PDFF and (not pdff.exists() or not pdwf.exists()): - water_nii = to_nii(water) - fat_nii = to_nii(fat) - water_nii.set_dtype_() - fat_nii.set_dtype_() - if not pdff.exists(): - nii = fat_nii / (water_nii + fat_nii) - nii[water_nii + fat_nii == 0] = 0 - nii *= 1000 - nii.set_dtype_("smallest_int") - nii.save(pdff) - if not pdwf.exists(): - nii = water_nii / (water_nii + fat_nii) - nii[water_nii + fat_nii == 0] = 0 - nii *= 1000 - nii.set_dtype_("smallest_int") - nii.save(pdwf) + if fat is not None and water is not None: + water_nii = to_nii(water) + fat_nii = to_nii(fat) + water_nii.set_dtype_() + fat_nii.set_dtype_() + if not pdff.exists(): + nii = fat_nii / (water_nii + fat_nii) + nii[water_nii + fat_nii == 0] = 0 + nii *= 1000 + nii.set_dtype_("smallest_int") + nii.save(pdff) + if not pdwf.exists(): + nii = water_nii / (water_nii + fat_nii) + nii[water_nii + fat_nii == 0] = 0 + nii *= 1000 + nii.set_dtype_("smallest_int") + nii.save(pdwf) + else: + # No raw fat/water — derive the missing fraction from the one the scanner shipped. + if not pdff.exists() and "mevibe_part-water-fraction" in fam: + wf_nii = to_nii(_check(fam["mevibe_part-water-fraction"])) + nii = 1000 - wf_nii + nii.set_dtype_("smallest_int") + nii.save(pdff) + if not pdwf.exists() and "mevibe_part-fat-fraction" in fam: + ff_nii = to_nii(_check(fam["mevibe_part-fat-fraction"])) + nii = 1000 - ff_nii + nii.set_dtype_("smallest_int") + nii.save(pdwf) + + # Downstream (see _apply_corrections_to_subj_dict) expects "mevibe_part-fat" to hold PDFF. if pdff.exists(): out["mevibe_part-fat"] = pdff - if pdff.exists(): + if pdwf.exists(): out["mevibe_part-fat"] = pdwf # else: # pdff = _check(fam["mevibe_part-fat-fraction_desc-reconstructed"]) @@ -567,10 +592,11 @@ def loop_over_repaired_nako( test=False, verbose=False, sort=True, - test_key="/100/10", # path matching. if you want on specific us a 6 digits + test_key="/102/", # path matching. if you want on specific us a 6 digits decision_cache: DecisionCache | Path | str | None = None, corrected_index: dict | Path | str | None = None, skip_subject=None, + vibe_mismatch_snap_dir: Path | str | None = None, ): """Iterate over the repaired NAKO dataset yielding per-subject file dicts. @@ -593,6 +619,12 @@ def loop_over_repaired_nako( test_key: Path substring passed to the BIDS scanner's ``filter_file`` when ``test=True``; only paths containing this substring are indexed. Defaults to a hard-coded example subject. baseline_metadata: Path to the NAKO baseline CSV used to look up height metadata. + vibe_mismatch_snap_dir: When set, subjects whose VIBE parts don't share a shape are skipped + (not yielded); a review snapshot is written to this directory as + ``sub-_vibe-shape-mismatch.jpg`` when that file doesn't already exist. + Two sub-folders ``accept/`` and ``reject/`` are also created on demand: if the reviewer + moves the jpg into ``accept/`` the next run auto-answers the grid prompt with "y" + (resample); if moved into ``reject/`` it auto-answers "n" (drop mismatched keys). Yields: Dict mapping short keys to ``BIDS_FILE`` entries for one subject. @@ -617,7 +649,6 @@ def loop_over_repaired_nako( "derivatives-fullbody-poi", # fullbody / fov101 / fov102 POIs + registered segmentations on the stitched-water grid ], filter_file=(lambda x: test_key in str(x)) if test else None, - ) for sub, subj in gbi.enumerate_subjects(sort=sort, shuffle=not sort): @@ -702,6 +733,8 @@ def loop_over_repaired_nako( q.filter_format("mevibe") # q.filter("sequ", "me1") mevibe_fams = list(q.loop_dict(key_addendum=["mod", "part", "desc"])) + # Drop derivative-only families that don't carry the raw echo images. + mevibe_fams = [f for f in mevibe_fams if "mevibe_part-eco0-opp1" in f] if len(mevibe_fams) > 1: labels = [str(f.get("mevibe_part-eco0-opp1", f)) for f in mevibe_fams] cached_pick = _cached_pick(cache.get(sub, "mevibe_fam")) @@ -791,6 +824,7 @@ def loop_over_repaired_nako( else: cache.set(sub, "vibe_fam", {"pick": labels[choice], "reason": reason}) vibe_fams = [vibe_fams[choice]] + _vibe_skip_subject = False for fam in vibe_fams: vibe_by_key: dict = {} for _k in ( @@ -803,6 +837,34 @@ def loop_over_repaired_nako( ): if fam.get(_k): vibe_by_key[_k] = fam[_k][0] + if vibe_mismatch_snap_dir is not None: + sigs = _vibe_grid_sigs(vibe_by_key) + if len(set(sigs.values())) > 1: + snap_dir = Path(vibe_mismatch_snap_dir) + accept_dir = snap_dir / "accept" + reject_dir = snap_dir / "reject" + for d in (snap_dir, accept_dir, reject_dir): + d.mkdir(parents=True, exist_ok=True) + snap_name = f"sub-{sub}_vibe-shape-mismatch.jpg" + snap_path = snap_dir / snap_name + if (accept_dir / snap_name).exists(): + # User verified this mismatch is fine — auto-answer "y" (resample). + cache.set(sub, "grid_mismatch:vibe", {"decision": "resample", "reason": "accepted-via-snap"}) + elif (reject_dir / snap_name).exists(): + # User rejected this subject's VIBE — auto-answer "n" (drop mismatched keys). + cache.set(sub, "grid_mismatch:vibe", {"decision": "remove", "reason": "rejected-via-snap"}) + else: + if not snap_path.exists(): + try: + _save_vibe_shape_mismatch_snapshot(vibe_by_key, snap_path) + except Exception as e: # noqa: BLE001 + log.on_warning(f"sub-{sub}: failed to save vibe mismatch snapshot: {e}") + log.on_warning( + f"sub-{sub}: VIBE grid mismatch {sigs} — awaiting review " + f"(move {snap_name} into accept/ or reject/); skipping subject" + ) + _vibe_skip_subject = True + break check_same_grid(cache, sub, "vibe", vibe_by_key, inphase_key="vibe_part-inphase") # Propagate the check's outcome back to ``fam`` so the downstream unpack # loop below picks up resampled files (or skips removed keys). @@ -834,6 +896,8 @@ def loop_over_repaired_nako( for k, k2 in mapp.items(): if k in fam: subj_dict[k2] = fam[k][0] + if _vibe_skip_subject: + continue vert, spine, poi = get_current_best_T2w_seg(sub) subj_dict["vert"] = vert subj_dict["spine"] = spine @@ -864,6 +928,69 @@ def _is_grid_only_json(path: Path | str) -> bool: _CANONICAL_DONE_ROOT = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical/.hardlink_done") +_VIBE_MISMATCH_SNAP_DIR = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako-canonical/snaps/vibe-missmatch") + + +def _vibe_grid_sigs(vibe_by_key: dict) -> dict[str, str]: + """Return ``{key: grid-signature-string}`` for each present VIBE part. + + Mirrors ``check_same_grid``'s comparison: reads ``bf.get_grid_info()`` and + stringifies it, skipping msk entries and files that have no NIfTI. Two + signatures being unequal is exactly what would cause ``check_same_grid`` + to prompt the user. + """ + sigs: dict[str, str] = {} + for k, bf in vibe_by_key.items(): + if bf is None: + continue + if getattr(bf, "format", None) == "msk" or k.startswith("msk"): + continue + get_nii_file = getattr(bf, "get_nii_file", None) + if callable(get_nii_file) and get_nii_file() is None: + continue + try: + g = bf.get_grid_info() + except Exception: # noqa: BLE001 + continue + if g is None: + continue + sigs[k] = str(g) + return sigs + + +def _save_vibe_shape_mismatch_snapshot(vibe_by_key: dict, out_path: Path) -> None: + """Save a review snapshot of VIBE parts whose grids don't agree. + + Renders one sagittal+coronal frame per present VIBE part (in/out/water/fat), + titled with the part's shape, so a reviewer can eyeball what's off. + """ + from TPTBox.spine.snapshot2D import Snapshot_Frame, create_snapshot + + frames = [] + for k in ("vibe_part-inphase", "vibe_part-outphase", "vibe_part-water", "vibe_part-fat"): + bf = vibe_by_key.get(k) + if bf is None: + continue + try: + shape = tuple(bf.get_grid_info().shape) # type: ignore[union-attr] + except Exception: # noqa: BLE001 + shape = None + frames.append( + Snapshot_Frame( + image=bf, + mode="MRI", + sagittal=True, + coronal=True, + axial=False, + crop_msk=False, + title=f"{k} shape={shape}", + ) + ) + if not frames: + return + out_path.parent.mkdir(parents=True, exist_ok=True) + create_snapshot(snp_path=[out_path], frames=frames) + def _hard_link_done_marker(sub: str) -> Path: return _CANONICAL_DONE_ROOT / f"{sub}.done" @@ -1403,6 +1530,17 @@ def _drain(fs): help="Trace hard_link()'s planned links for one subject; check src exists + same fs as target.", ) parser.add_argument("--workers", type=int, default=None, help="Number of worker processes (default: cpu_count-1).") + parser.add_argument( + "--vibe-mismatch-snaps", + nargs="?", + const=str(_VIBE_MISMATCH_SNAP_DIR), + # default=None, + default=str(_VIBE_MISMATCH_SNAP_DIR), + metavar="DIR", + help="Skip subjects whose VIBE parts have grid mismatches; write a review .jpg to DIR " + f"(defaults to {_VIBE_MISMATCH_SNAP_DIR}) unless one already exists there. " + "Pass '' to disable.", + ) args = parser.parse_args() test = True @@ -1414,5 +1552,10 @@ def _drain(fs): precompute_grid_info_parallel(num_workers=args.workers, test=test) else: corrected = load_corrected_index() - for d in loop_over_repaired_nako(test=test, corrected_index=corrected, skip_subject=is_hard_linked): + for d in loop_over_repaired_nako( + test=test, + corrected_index=corrected, + skip_subject=is_hard_linked, + vibe_mismatch_snap_dir=args.vibe_mismatch_snaps or None, + ): hard_link(d) From afe83299f16d81346b982c60682d738dee3d189a Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Thu, 17 Sep 2026 11:59:56 +0000 Subject: [PATCH 12/26] add versioning --- TPTBox/spine/spinestats/_run_all.py | 93 +++++++++++++++++++++++++++++ 1 file changed, 93 insertions(+) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 236d189b..bbd642de 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -135,6 +135,9 @@ def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple - json fails to parse, - any required main key is missing from the loaded dict. """ + # TODO(hash-invalidation): also compare each entry in data["_provenance"]["inputs"] + # against the current file's sha1 (see _build_provenance) and force recompute on + # mismatch. Deferred while we assume inputs are unchanged. if not json_path.exists(): return False, None json_mtime = json_path.stat().st_mtime @@ -153,6 +156,93 @@ def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple return True, data +_PROVENANCE_INPUT_KEYS: tuple[str, ...] = ( + "t2w", + "vibe_part-water", + "vibe_part-fat", + "vibe_part-inphase", + "vibe_part-outphase", + "vibe-water", + "vibe-fat", + "vibe-inphase", + "vibe-outphase", + "vert", + "spine", + "vibeseg100", + "roi", + "fullbody_poi", +) + + +def _resolve_prov_path(v) -> Path | None: + """Best-effort ``file_dict`` value → filesystem Path for provenance recording.""" + if v is None: + return None + if isinstance(v, BIDS_FILE): + nii = v.get_nii_file() + if nii is not None: + return Path(nii) + j = v.file.get("json") if hasattr(v, "file") else None + return Path(j) if j is not None else None + if isinstance(v, (str, Path)): + s = str(v) + return Path(s) if s else None + return None + + +def _file_provenance(path: Path, prior: dict | None = None) -> dict: + """Return provenance dict for ``path``. + + Shape is either ``{"path", "mtime_ns", "sha1"}`` or, when the file is + missing, ``{"path", "missing": True}``. When ``prior`` has the same + ``mtime_ns`` as the current file, its ``sha1`` is reused to avoid + re-hashing — this is what keeps reruns cheap while the "assume inputs + unchanged" mode is in effect. + """ + import hashlib + + p = Path(path) + if not p.exists(): + return {"path": str(p), "missing": True} + st = p.stat() + if isinstance(prior, dict) and prior.get("mtime_ns") == st.st_mtime_ns and isinstance(prior.get("sha1"), str): + return {"path": str(p), "mtime_ns": st.st_mtime_ns, "sha1": prior["sha1"]} + h = hashlib.sha1() + with p.open("rb") as f: + for chunk in iter(lambda: f.read(1 << 20), b""): + h.update(chunk) + return {"path": str(p), "mtime_ns": st.st_mtime_ns, "sha1": h.hexdigest()} + + +def _build_provenance(file_dict: dict, poi_out: Path | str | None, prior: dict | None) -> dict: + """Build the ``_provenance`` block for an aggregated per-subject JSON. + + Records ``path`` / ``mtime_ns`` / ``sha1`` for every known input in + ``file_dict`` plus the POI json at ``poi_out``. Reuses + ``prior["inputs"][k]`` to skip re-hashing files whose mtime is unchanged. + """ + from datetime import datetime, timezone + + prior_inputs = (prior or {}).get("inputs", {}) if isinstance(prior, dict) else {} + inputs: dict[str, dict] = {} + for key in _PROVENANCE_INPUT_KEYS: + if key not in file_dict: + continue + p = _resolve_prov_path(file_dict[key]) + if p is None: + continue + inputs[key] = _file_provenance(p, prior_inputs.get(key)) + if poi_out is not None: + p = Path(poi_out) + if p.exists(): + inputs["poi"] = _file_provenance(p, prior_inputs.get("poi")) + return { + "version": 1, + "written_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), + "inputs": inputs, + } + + def run_all( file_dict, override: bool = False, @@ -318,6 +408,7 @@ def _need(*keys: str, compute: bool) -> bool: if not (need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso or need_curvature or need_pelvic): if _merge_endplate_angles(out, Path(poi_out)) or save: + out["_provenance"] = _build_provenance(file_dict, poi_out, out.get("_provenance")) save_json(final_out, out) return out @@ -451,6 +542,7 @@ def _need(*keys: str, compute: bool) -> bool: logger.print_error() _merge_endplate_angles(out, Path(poi_out)) + out["_provenance"] = _build_provenance(file_dict, poi_out, out.get("_provenance")) logger.on_save("save", final_out.name) save_json(final_out, out) return out @@ -604,6 +696,7 @@ def _rows_from_json(subject_id: str, data: dict) -> tuple[dict[str, Any], list[d "axial_rotation", "endplate_internal_angle", "segmental_endplate_angles", + "_provenance", ) subject_view = {k: v for k, v in data.items() if k not in _PER_SUBJECT_EXCLUDE} _flatten("", subject_view, per_subject) From 946742914d301dda8287329a2c53d9178414d711 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Thu, 17 Sep 2026 12:56:25 +0000 Subject: [PATCH 13/26] add Verida lordosie version --- TPTBox/spine/spinestats/_load_nako.py | 3 + TPTBox/spine/spinestats/_load_nako_wh.py | 22 ++++ TPTBox/spine/spinestats/_run_all.py | 29 +++++- TPTBox/spine/spinestats/angles.py | 7 ++ TPTBox/spine/spinestats/veridah_angles.py | 118 ++++++++++++++++++++++ 5 files changed, 177 insertions(+), 2 deletions(-) create mode 100644 TPTBox/spine/spinestats/veridah_angles.py diff --git a/TPTBox/spine/spinestats/_load_nako.py b/TPTBox/spine/spinestats/_load_nako.py index b0462ee2..e2a98c64 100644 --- a/TPTBox/spine/spinestats/_load_nako.py +++ b/TPTBox/spine/spinestats/_load_nako.py @@ -7,6 +7,7 @@ from TPTBox import Print_Logger from TPTBox.core.bids_files import BIDS_FILE, BIDS_Family, Buffered_BIDS_Global_info from TPTBox.core.nii_wrapper import to_nii +from TPTBox.spine.spinestats._load_nako_wh import get_current_best_VERIDAH as _get_current_best_VERIDAH log = Print_Logger() @@ -320,6 +321,8 @@ def loop_over_repaired_nako( subj_dict["vert"] = vert subj_dict["spine"] = spine subj_dict["poi"] = poi + veridah = _get_current_best_VERIDAH(sub) + subj_dict["veridah"] = str(veridah) if veridah is not None else None yield subj_dict diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index 2d09d7ea..6cbb1a56 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -523,6 +523,25 @@ def get_corrected_mevibe(fam: BIDS_Family, compute_PDFF=True): # TODO return di return out +def get_current_best_VERIDAH(sub) -> Path | None: + """Return the newest VERIDAH-label JSON for ``sub`` (V2 preferred), or ``None`` if missing. + + The file lives alongside the T2w segmentation under + ``derivatives_spine_inference_162_sacrumfix///T2w/`` and holds the + ``orig_label -> fpath`` remapping used by :func:`compute_veridah_variants`. + """ + sub = str(sub).split("_")[0].replace("sub-", "") + for folder in ("derivatives_spine_inference_162_sacrumfix",): + for suffix in ("VERIDAH-label-V2", "VERIDAH-label"): + p = Path( + f"/DATA/NAS/datasets_processed/NAKO/dataset-nako/{folder}/{sub[:3]}/{sub}/T2w/" + f"sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-vert_desc-{suffix}_stat.json" + ) + if p.exists(): + return p + return None + + def get_current_best_T2w_seg(sub, black_list_t2w=None): if black_list_t2w is None: black_list_t2w = [ @@ -902,6 +921,8 @@ def loop_over_repaired_nako( subj_dict["vert"] = vert subj_dict["spine"] = spine subj_dict["poi"] = poi + veridah = get_current_best_VERIDAH(sub) + subj_dict["veridah"] = str(veridah) if veridah is not None else None if corrected_index: _apply_corrections_to_subj_dict(str(sub), subj_dict, corrected_index) verify_missing_images(cache, sub, subj_dict) @@ -1049,6 +1070,7 @@ def hard_link( "vert", # t2w (stiched) "spine", # t2w (stiched) "poi", # t2w (stiched) + "veridah", # VERIDAH enumeration-anomaly relabeling (V2 preferred) ] imgs = [ "pd", diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index bbd642de..75059029 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -96,6 +96,16 @@ def get_nako_paths(nako_id: str) -> dict[str, Path | None]: fullbody_poi = ( DATASET_ROOT / f"derivatives-fullbody-poi/{pfx}/{sub}/vibe/sub-{sub}_sequ-stitched_acq-ax_part-water_seg-fullbody_poi.json" ) + veridah = None + for suffix in ("VERIDAH-label-V2", "VERIDAH-label"): + p = ( + DATASET_ROOT + / f"derivatives_spine_inference_162_sacrumfix/{pfx}/{sub}/T2w/" + f"sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-vert_desc-{suffix}_stat.json" + ) + if p.exists(): + veridah = p + break roi = ( DATASET_ROOT / f"derivatives_Abdominal-Segmentation/{pfx}/{sub}/vibe/sub-{nako_id}_sequ-stitched_acq-ax_mod-vibe_seg-ROI_msk.nii.gz" ) @@ -112,6 +122,7 @@ def get_nako_paths(nako_id: str) -> dict[str, Path | None]: "roi": roi, "vibeseg100": vibeseg100 if vibeseg100.exists() else None, "fullbody_poi": fullbody_poi if fullbody_poi.exists() else None, + "veridah": veridah, "dataset": DATASET_ROOT, } @@ -171,6 +182,7 @@ def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple "vibeseg100", "roi", "fullbody_poi", + "veridah", ) @@ -398,7 +410,8 @@ def _need(*keys: str, compute: bool) -> bool: _merge_per_vertebra_metrics(out) save = True #### - need_poi = need_cobb or need_ivd or need_vert or need_curvature + need_veridah = override or "curv_veridah" not in out + need_poi = need_cobb or need_ivd or need_vert or need_curvature or need_veridah need_t2w = need_ivd or need_vert or need_vbq need_vert_nii = need_poi or need_vbq or need_bcs or need_mfi need_spine_nii = need_vert_nii or need_vbq @@ -406,7 +419,9 @@ def _need(*keys: str, compute: bool) -> bool: need_roi = need_mfi or need_torso need_vibe_wf = need_mfi - if not (need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso or need_curvature or need_pelvic): + if not ( + need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso or need_curvature or need_pelvic or need_veridah + ): if _merge_endplate_angles(out, Path(poi_out)) or save: out["_provenance"] = _build_provenance(file_dict, poi_out, out.get("_provenance")) save_json(final_out, out) @@ -495,6 +510,16 @@ def _need(*keys: str, compute: bool) -> bool: logger.on_fail("curvature error caught") logger.print_error() + if need_veridah and poi is not None: + try: + from TPTBox.spine.spinestats.veridah_angles import compute_veridah_variants + + logger.on_debug("veridah variants") + out["curv_veridah"] = compute_veridah_variants(poi, file_dict.get("veridah")) + except Exception: + logger.on_fail("veridah variants error caught") + logger.print_error() + # Merge wedge metrics directly into the per-label vert_geometry / ivd_geometry # entries so they land in per_vertebra.xlsx / per_ivd.xlsx automatically. for geom_key in ("vert_geometry", "ivd_geometry"): diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 606f56a5..122a24e8 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -686,6 +686,7 @@ def plot_compute_lordosis_and_kyphosis( seg: Image_Reference | None = None, line_len=100, project_2D=True, + curvature_definition=curvature_definition, ) -> tuple[dict[str, float | None], Snapshot_Frame]: """Plots and computes the angles of lordosis and kyphosis on a spinal image. @@ -701,6 +702,12 @@ def plot_compute_lordosis_and_kyphosis( seg (Image_Reference | None): The segmentation image reference. Optional, can be None. line_len (int): The length of the lines representing the vertebrae directions (default is 100). project_2D (bool, optional): If True, the angles are computed in the 2D sagittal projection; otherwise in 3D. Defaults to True. + curvature_definition (dict[str, Def_Curvature], optional): Mapping of output-key name + → :class:`Def_Curvature` describing which vertebra pair defines each angle. + Defaults to the module-level ``curvature_definition`` (cervical_lordosis, + thoracic_kyphosis, lumbar_lordosis). Pass a custom dict to compute a different + set of segmental angles or to override the ``last_thoracic`` / ``last_lumbar`` + resolution — the output dict's keys mirror this mapping's keys. Returns: tuple: A tuple containing: diff --git a/TPTBox/spine/spinestats/veridah_angles.py b/TPTBox/spine/spinestats/veridah_angles.py new file mode 100644 index 00000000..2c40bb54 --- /dev/null +++ b/TPTBox/spine/spinestats/veridah_angles.py @@ -0,0 +1,118 @@ +"""VERIDAH-aware lordosis / kyphosis variants. + +Two extra angle computations on top of the standard ``curv`` block in +:mod:`TPTBox.spine.spinestats.angles`: + +- ``anomaly``: relabel the POI using the VERIDAH ``orig_label -> fpath`` + map (so a supernumerary T13 lands on label 28 instead of shifting the + entire lumbar enumeration down by one) and recompute all three + regional angles (cervical / thoracic / lumbar). +- ``k4 / k5 / k6 / k7``: force the last ``k`` present vertebrae above + the sacrum to be interpreted as L1..Lk, drop everything cranial, and + compute **only** ``lumbar_lordosis``. Independent of VERIDAH. + +Both variants live under the top-level JSON key ``curv_veridah`` (see +``_run_all.py``); the standard ``curv`` output is untouched. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +from TPTBox.core.poi import POI, POI_Descriptor +from TPTBox.core.vert_constants import Vertebra_Instance +from TPTBox.spine.spinestats.angles import compute_lordosis_and_kyphosis + + +def _load_veridah(path: Path) -> dict | None: + """Return the single dict inside a VERIDAH ``_stat.json``, or ``None`` on error.""" + try: + data = json.loads(Path(path).read_text()) + except (OSError, json.JSONDecodeError): + return None + if isinstance(data, list) and data and isinstance(data[0], dict): + return data[0] + if isinstance(data, dict): + return data + return None + + +def _relabel_poi(poi: POI, mapping: dict[int, int]) -> POI: + """Return a copy of ``poi`` whose region ids are re-keyed via ``mapping``. + + Regions absent from ``mapping`` are dropped. Non-vertebra regions + (i.e. those already outside the vertebra label range) fall through + with their original id — but in practice a POI produced by + :func:`calc_poi_from_subreg_vert` only carries vertebra regions. + """ + new_centroids = POI_Descriptor() + for region, subregion, coord in poi.centroids.items(): + if region in mapping: + new_centroids[(mapping[region], subregion)] = coord + return poi.copy(centroids=new_centroids) + + +def _last_k_relabel(poi: POI, k: int) -> POI: + """Keep the ``k`` most caudal non-sacral vertebrae; relabel them L1..Lk. + + Uses :meth:`Vertebra_Instance.order` to walk cranio-caudal, ignores + sacrum members, and keeps the last ``k`` present labels. The bottom + one becomes ``Lk`` (label ``20 + k - 1``), stepping up by one to + ``L1`` (label 20). + """ + order = Vertebra_Instance.order() + sacrum_vals = {v.value for v in Vertebra_Instance.sacrum()} + present = {r for r, _s, _c in poi.centroids.items()} + non_sacral_in_order = [v.value for v in order if v.value in present and v.value not in sacrum_vals] + if len(non_sacral_in_order) < k: + return poi.copy(centroids=POI_Descriptor()) + last_k = non_sacral_in_order[-k:] # cranio → caudal + # last_k[0] should be L1 (20), last_k[-1] should be Lk (20+k-1). + mapping = {orig: 20 + i for i, orig in enumerate(last_k)} + return _relabel_poi(poi, mapping) + + +def _compute_anomaly_variant(poi: POI, veridah_json_path: Path | None) -> dict[str, float | None] | None: + """Recompute the three regional angles after applying the VERIDAH label correction. + + Returns ``None`` when the VERIDAH file is missing or unusable. + """ + if veridah_json_path is None: + return None + veridah = _load_veridah(Path(veridah_json_path)) + if veridah is None: + return None + orig = veridah.get("orig_label") + fpath = veridah.get("fpath") + if not isinstance(orig, list) or not isinstance(fpath, list) or len(orig) != len(fpath): + return None + mapping = {int(o): int(f) for o, f in zip(orig, fpath)} + return compute_lordosis_and_kyphosis(_relabel_poi(poi, mapping)) + + +def _compute_k_variant(poi: POI, k: int) -> dict[str, float | None]: + """Compute lumbar-lordosis only, with the last ``k`` vertebrae treated as L1..Lk.""" + relabeled = _last_k_relabel(poi, k) + if not relabeled.centroids: + return {"lumbar_lordosis": None} + full = compute_lordosis_and_kyphosis(relabeled) + return {"lumbar_lordosis": full.get("lumbar_lordosis")} + + +def compute_veridah_variants(poi: POI, veridah_json_path: Path | str | None) -> dict: + """Return the ``curv_veridah`` block for one subject. + + Shape:: + + { + "anomaly": {"cervical_lordosis", "thoracic_kyphosis", "lumbar_lordosis"} | None, + "k4": {"lumbar_lordosis": …}, + "k5": {...}, "k6": {...}, "k7": {...}, + } + """ + veridah_path = Path(veridah_json_path) if veridah_json_path is not None else None + out: dict = {"anomaly": _compute_anomaly_variant(poi, veridah_path)} + for k in (4, 5, 6, 7): + out[f"k{k}"] = _compute_k_variant(poi, k) + return out From 06a5f1348fadf9d8c056f1b24b88662766a48d8d Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 12:26:46 +0000 Subject: [PATCH 14/26] add S1 Endplate --- TPTBox/spine/spinestats/poi_fun/endplates.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/TPTBox/spine/spinestats/poi_fun/endplates.py b/TPTBox/spine/spinestats/poi_fun/endplates.py index 48c87302..fde98bba 100644 --- a/TPTBox/spine/spinestats/poi_fun/endplates.py +++ b/TPTBox/spine/spinestats/poi_fun/endplates.py @@ -508,7 +508,7 @@ def endplate_to_super_infer_endplate(vert: NII, spine: NII) -> tuple[NII, NII]: spine = spine.copy() vert_org = vert.copy() vert[vert >= 40] = 0 - vert[spine.extract_label([Location.Vertebra_Corpus, Location.Vertebra_Corpus_border]) != 1] = 0 + vert[spine.extract_label([Location.Vertebra_Corpus, Location.Vertebra_Corpus_border, Vertebra_Instance.S1]) != 1] = 0 vert %= 100 v = vert.infect( spine.extract_label( @@ -522,8 +522,10 @@ def endplate_to_super_infer_endplate(vert: NII, spine: NII) -> tuple[NII, NII]: verbose=False, ) endplate_nii = v * endplate_nii + spine[endplate_nii == Vertebra_Instance.S1.value] = Location.Sacrum_Endplate.value spine[np.logical_and(endplate_nii == vert_org % 100, endplate_nii != 0)] = Location.Vertebral_Body_Endplate_Inferior.value spine[spine == Location.Endplate.value] = Location.Vertebral_Body_Endplate_Superior.value + vert_org[endplate_nii != 0] = v[endplate_nii != 0] + 200 return vert_org, spine From 840e3534d378746b382ee045884d91cfdb726193 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 12:27:09 +0000 Subject: [PATCH 15/26] add warning to not use is_ct in "DAExt" --- TPTBox/core/internal/train_nnUnet/prepere_dataset.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/TPTBox/core/internal/train_nnUnet/prepere_dataset.py b/TPTBox/core/internal/train_nnUnet/prepere_dataset.py index bf9a4da0..ced21f45 100644 --- a/TPTBox/core/internal/train_nnUnet/prepere_dataset.py +++ b/TPTBox/core/internal/train_nnUnet/prepere_dataset.py @@ -199,6 +199,13 @@ def _validate_config(cfg: DatasetConfig) -> None: "Use 'nnUNetTrainerNoMirroring', a SmaugLab DAExt trainer (mirror is stripped from its " "params JSON), or drop the mirror pairs." ) + # SmaugLab DAExt trainers assume the SmaugLab (non-CT) preprocessing path; combining them with + # is_ct=True mixes CT normalization with augmentations tuned for the SmaugLab regime. + if cfg.is_ct and "DAExt" in cfg.nn_trainer: + logger.on_warning( + f"is_ct=True combined with nn_trainer='{cfg.nn_trainer}': SmaugLab DAExt trainers are " + "not intended to be used with CT normalization. Set is_ct=False or pick a non-DAExt trainer." + ) if errors: logger.on_fail("Config validation failed:") for e in errors: From b8e6f3c525a350c7e1b15987f0d0359735555710 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 13:05:25 +0000 Subject: [PATCH 16/26] Auto-transform of direction-vector fields in POI.info Adds a small registration mechanism so producers can attach per-vertebra direction vectors (endplate PCA normals, orientation vectors, ...) or label-keyed scalars (endplate_internal_angle, curvature_*) to poi.info and have them stay aligned with the POI points automatically: - New TPTBox/core/poi_fun/vector_fields.py with the constants POI_INFO_VECTOR_FIELDS_KEY / POI_INFO_LABEL_KEYED_FIELDS_KEY and three helpers (_transform_direction_vectors_inplace, _rotate_direction_vectors_inplace, _remap_vector_field_keys_inplace). - POI.reorient signed-permutes vectors; POI.resample_from_to applies the real R_ref.T @ R_self rotation; POI.rescale is a no-op. - POI_Global.to_cord_system flips vector components on RAS<->LPS. - POI.map_labels remaps top-level keys of both registries. - endplates.py registers angle_superior/inferior_endplate, endplate_internal_angle, and curvature_* when they're produced. - Fix: vertebra_pois_non_centroids.py now (re)runs calc_endplate_points_ when S1 is present but its Sacrum_Endplate- derived superior-endplate landmark is missing, so the sacrum endplate direction actually lands in the buffer. - Angles: _endplate_ap_direction preferred over Vertebra_Direction_ Posterior for MoveTo.TOP/BOTTOM in _get_norm, yielding endplate- aligned Cobb-style lordosis/kyphosis lines. Docs: new section in docs/api/poi_fun.md. Tests: unit_tests/test_poi_vector_fields.py (13 tests). Co-Authored-By: Claude Opus 4.7 --- TPTBox/core/poi.py | 30 ++- TPTBox/core/poi_fun/poi_abstract.py | 11 +- TPTBox/core/poi_fun/poi_global.py | 10 +- TPTBox/core/poi_fun/vector_fields.py | 133 ++++++++++++ .../poi_fun/vertebra_pois_non_centroids.py | 7 +- TPTBox/spine/spinestats/angles.py | 71 ++++++- TPTBox/spine/spinestats/poi_fun/endplates.py | 16 ++ docs/api/poi_fun.md | 63 ++++++ unit_tests/test_poi_vector_fields.py | 200 ++++++++++++++++++ 9 files changed, 527 insertions(+), 14 deletions(-) create mode 100644 TPTBox/core/poi_fun/vector_fields.py create mode 100644 unit_tests/test_poi_vector_fields.py diff --git a/TPTBox/core/poi.py b/TPTBox/core/poi.py index eaf38b19..7365fdff 100755 --- a/TPTBox/core/poi.py +++ b/TPTBox/core/poi.py @@ -45,6 +45,16 @@ ### CURRENT TYPE DEFINITIONS C = TypeVar("C", bound="POI") + +# Direction-vector auto-transform helpers live in poi_fun.vector_fields. +# Re-exported here for backwards compatibility with existing callers. +from TPTBox.core.poi_fun.vector_fields import ( # noqa: E402 + POI_INFO_VECTOR_FIELDS_KEY, + _remap_vector_field_keys_inplace, + _rotate_direction_vectors_inplace, + _transform_direction_vectors_inplace, +) + POI_Reference = Union[ bids_files.BIDS_FILE, Path, @@ -522,8 +532,11 @@ def reorient( self.shape = shape self.origin = origin self.rotation = rotation + _transform_direction_vectors_inplace(self.info, trans) return self - return self.copy(orientation=axcodes_to, centroids=points, zoom=zoom, shape=shape, origin=origin, rotation=rotation) + new_poi = self.copy(orientation=axcodes_to, centroids=points, zoom=zoom, shape=shape, origin=origin, rotation=rotation) + _transform_direction_vectors_inplace(new_poi.info, trans) + return new_poi def reorient_(self, axcodes_to: AX_CODES = ("P", "I", "R"), decimals=3, verbose: logging = False, _shape=None) -> Self: """In-place variant of :meth:`reorient`.""" @@ -532,6 +545,11 @@ def reorient_(self, axcodes_to: AX_CODES = ("P", "I", "R"), decimals=3, verbose: def rescale(self, voxel_spacing: ZOOMS = (1, 1, 1), decimals=ROUNDING_LVL, verbose: logging = True, inplace=False) -> Self: """Rescale the POI coordinates to a new voxel spacing in the current x-y-z-orientation. + Direction-vector fields registered under ``info[POI_INFO_VECTOR_FIELDS_KEY]`` + are left untouched: they are assumed to live in mm-space aligned with the + current voxel axes (see :func:`_transform_direction_vectors_inplace`), and + a spacing change relabels axes without rotating them. + Args: voxel_spacing (tuple[float, float, float], optional): New voxel spacing in millimeters. Defaults to (1, 1, 1). decimals (int, optional): Number of decimal places to round the rescaled coordinates to. Defaults to ROUNDING_LVL. @@ -612,7 +630,11 @@ def to_global(self, itk_coords=False) -> POI_Global: ) def resample_from_to(self, ref: Has_Grid) -> POI: - """Resample this POI to the grid of another image by converting to global and back. + """Resample this POI to the voxel grid of ``ref``. + + Registered direction-vector fields (see :data:`POI_INFO_VECTOR_FIELDS_KEY`) + are rotated by ``R_ref^T @ R_self`` so their components stay aligned with + the target grid's axes. Args: ref (Has_Grid): Target image grid (any object providing affine/orientation info). @@ -620,7 +642,9 @@ def resample_from_to(self, ref: Has_Grid) -> POI: Returns: POI: A new POI in the voxel space of ``ref``. """ - return self.to_global().to_other(ref) + out = self.to_global().to_other(ref) + _rotate_direction_vectors_inplace(out.info, self.rotation, getattr(ref, "rotation", None)) + return out def resample_from_to_(self, ref: Has_Grid) -> Self: """In-place variant of :meth:`resample_from_to`.""" diff --git a/TPTBox/core/poi_fun/poi_abstract.py b/TPTBox/core/poi_fun/poi_abstract.py index d0151a19..97789148 100755 --- a/TPTBox/core/poi_fun/poi_abstract.py +++ b/TPTBox/core/poi_fun/poi_abstract.py @@ -558,12 +558,19 @@ def map_labels( continue poi_new[region:subreg] = value new_values = poi_new + from TPTBox.core.poi_fun.vector_fields import _remap_vector_field_keys_inplace + if new_values is None: - return self if inplace else self.copy() + out = self if inplace else self.copy() + _remap_vector_field_keys_inplace(out.info, label_map_region_) + return out if inplace: self.centroids = new_values + _remap_vector_field_keys_inplace(self.info, label_map_region_) return self - return self.copy(centroids=new_values) + out = self.copy(centroids=new_values) + _remap_vector_field_keys_inplace(out.info, label_map_region_) + return out def map_labels_( self, diff --git a/TPTBox/core/poi_fun/poi_global.py b/TPTBox/core/poi_fun/poi_global.py index 016e1fdc..c9797512 100755 --- a/TPTBox/core/poi_fun/poi_global.py +++ b/TPTBox/core/poi_fun/poi_global.py @@ -158,7 +158,8 @@ def to_cord_system(self, itk_coords: bool, inplace: bool = False) -> Self: """Convert between ITK (LPS) and NIfTI (RAS) coordinate systems. Flips the first two coordinate axes when switching between the two - systems (LPS ↔ RAS only differs in the sign of x and y). + systems (LPS ↔ RAS only differs in the sign of x and y). Registered + direction-vector fields in ``info`` are flipped along the same axes. Args: itk_coords: ``True`` for ITK/LPS output, ``False`` for NIfTI/RAS. @@ -167,12 +168,19 @@ def to_cord_system(self, itk_coords: bool, inplace: bool = False) -> Self: Returns: ``POI_Global`` in the requested coordinate system. """ + import numpy as np + + from TPTBox.core.poi_fun.vector_fields import _transform_direction_vectors_inplace + out = self if inplace else self.copy() if self.itk_coords == itk_coords: return out out.itk_coords = itk_coords for k1, k2, v in self.items(): out[k1, k2] = (-v[0], -v[1], v[2]) + # x and y are negated, z is kept -- express as a signed-permutation trans. + trans = np.array([[0, -1], [1, -1], [2, 1]], dtype=int) + _transform_direction_vectors_inplace(out.info, trans) return out def to_other(self, msk: Has_Grid, verbose=False) -> poi.POI: diff --git a/TPTBox/core/poi_fun/vector_fields.py b/TPTBox/core/poi_fun/vector_fields.py new file mode 100644 index 00000000..30657040 --- /dev/null +++ b/TPTBox/core/poi_fun/vector_fields.py @@ -0,0 +1,133 @@ +"""Helpers for auto-transforming direction-vector fields stored in ``POI.info``. + +Producers write direction vectors into ``poi.info[""]`` as +``{key: (x, y, z)}`` dicts and register ```` in +``poi.info[POI_INFO_VECTOR_FIELDS_KEY]``. The helpers here are then called by +``POI.reorient`` / ``POI.resample_from_to`` / ``POI_Global.to_cord_system`` / +``POI.map_labels`` to keep the vectors and their keys aligned with the POI +points as the POI is transformed. + +Vectors are assumed to live in mm-space aligned with the POI's current voxel +axes; ``rescale`` is therefore a no-op for them. +""" + +from __future__ import annotations + +import numpy as np + +from TPTBox.core.vert_constants import Vertebra_Instance + +# poi.info key naming the info dicts whose values are direction vectors +# in the POI's current axis frame (i.e. mm-space aligned with poi.orientation). +POI_INFO_VECTOR_FIELDS_KEY = "_vector_fields" + +# poi.info key naming *additional* info dicts (typically scalar-valued, e.g. +# per-vertebra angles or curvature) whose top-level keys are region labels or +# ``Vertebra_Instance`` names. They participate in :func:`map_labels` key +# remapping but not in reorient / resample / to_cord_system vector transforms. +POI_INFO_LABEL_KEYED_FIELDS_KEY = "_label_keyed_fields" + + +def _transform_direction_vectors_inplace(info: dict, trans: np.ndarray) -> None: + """Reorient every registered direction-vector field in ``info`` according to ``trans``. + + ``trans`` is the output of ``nibabel.orientations.ornt_transform`` applied to + the POI's source and target axcodes. Each value stored under a registered + field must be a 3-tuple / 3-list (or ``None`` / non-3-tuple, which is skipped) + interpreted as a unit direction in the POI's current mm-space voxel-axis + frame. Values are updated in place. + """ + field_names = info.get(POI_INFO_VECTOR_FIELDS_KEY) + if not field_names: + return + perm = np.asarray(trans[:, 0], dtype=int) + flip = np.asarray(trans[:, 1], dtype=int) + for name in field_names: + vectors = info.get(name) + if not isinstance(vectors, dict): + continue + for k, v in list(vectors.items()): + if v is None: + continue + try: + v_arr = np.asarray(v, dtype=float) + except (TypeError, ValueError): + continue + if v_arr.shape != (3,): + continue + new_v = np.zeros(3, dtype=float) + new_v[perm] = v_arr * flip + vectors[k] = tuple(float(x) for x in new_v) + + +def _remap_vector_field_keys_inplace(info: dict, region_map: dict) -> None: + """Remap the top-level keys of every registered label-keyed field via ``region_map``. + + Considers fields registered under both :data:`POI_INFO_VECTOR_FIELDS_KEY` + (direction vectors) and :data:`POI_INFO_LABEL_KEYED_FIELDS_KEY` (scalar + per-label fields like ``endplate_internal_angle`` or ``curvature_*``). + Keys may be either integer region labels or ``Vertebra_Instance``-name + strings ("L1", "T12", ...); both are matched against ``region_map`` + (int-keyed). Duplicates after remapping keep the last write. No-op if + no fields are registered or ``region_map`` is empty. + """ + if not region_map: + return + field_names: list[str] = [] + for key in (POI_INFO_VECTOR_FIELDS_KEY, POI_INFO_LABEL_KEYED_FIELDS_KEY): + names = info.get(key) + if names: + field_names.extend(names) + if not field_names: + return + for name in field_names: + vectors = info.get(name) + if not isinstance(vectors, dict): + continue + remapped = {} + for k, v in vectors.items(): + new_k = k + if isinstance(k, int) and k in region_map: + new_k = region_map[k] + elif isinstance(k, str): + try: + label = Vertebra_Instance[k].value + except KeyError: + label = None + if label is not None and label in region_map: + try: + new_k = Vertebra_Instance(region_map[label]).name + except ValueError: + new_k = k + remapped[new_k] = v + vectors.clear() + vectors.update(remapped) + + +def _rotate_direction_vectors_inplace(info: dict, src_rot, tgt_rot) -> None: + """Rotate registered direction-vector fields from ``src_rot`` to ``tgt_rot``. + + Composes ``tgt_rot.T @ src_rot`` and applies it in place. No-op when either + rotation is None or no vector fields are registered. + """ + if src_rot is None or tgt_rot is None: + return + field_names = info.get(POI_INFO_VECTOR_FIELDS_KEY) + if not field_names: + return + R = np.asarray(tgt_rot, dtype=float).T @ np.asarray(src_rot, dtype=float) + for name in field_names: + vectors = info.get(name) + if not isinstance(vectors, dict): + continue + for k, v in list(vectors.items()): + if v is None: + continue + try: + v_arr = np.asarray(v, dtype=float) + except (TypeError, ValueError): + continue + if v_arr.shape != (3,): + continue + new_v = R @ v_arr + vectors[k] = tuple(float(x) for x in new_v) diff --git a/TPTBox/core/poi_fun/vertebra_pois_non_centroids.py b/TPTBox/core/poi_fun/vertebra_pois_non_centroids.py index 8f0583a9..beb4d04e 100755 --- a/TPTBox/core/poi_fun/vertebra_pois_non_centroids.py +++ b/TPTBox/core/poi_fun/vertebra_pois_non_centroids.py @@ -365,7 +365,12 @@ def compute_non_centroid_pois( # noqa: C901 log.on_text("Compute Vertebra Endplate DIRECTIONS", verbose=verbose) sub_regions = poi.keys_subregion() - if any(a.value not in sub_regions for a in endplate[:2]): # skip if all exists + # Also (re)run when S1 is present in the segmentation but its Sacrum_Endplate-derived + # superior-endplate landmark is not in the POI yet -- otherwise the sacrum block inside + # calc_endplate_points_ never gets a chance to populate poi.info["angle_superior_endplate"]["S1"]. + s1 = Vertebra_Instance.S1.value + sacrum_endplate_missing = s1 in _vert_ids and (s1, Location.Vertebral_Body_Endplate_Superior.value) not in poi + if any(a.value not in sub_regions for a in endplate[:2]) or sacrum_endplate_missing: # skip if all exists poi, *_ = calc_endplate_points_(poi, vert, subreg, _vert_ids=_vert_ids, log=log) ### STEP 1 Vert Direction### if Location.Vertebra_Direction_Inferior in locations: diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 122a24e8..f5e41c66 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -66,10 +66,10 @@ def get_location(self, v: Vertebra_Instance | int, poi: POI) -> tuple: subreg = Location.Additional_Vertebral_Body_Middle_Inferior_Median if (v, subreg) in poi: return (v, subreg) - # Test if it has next POINT + # Fall back to averaging v's centroid with the next vertebra's centroid. next_vert = v.get_next_poi(poi) - if next_vert is not None and (next_vert, 50) in poi: - return (v, subreg, next_vert, subreg) + if next_vert is not None and (v, 50) in poi and (next_vert, 50) in poi: + return (v, 50, next_vert, 50) elif self == self.TOP: prev_vert = v.get_previous_poi(poi) # Test IVD @@ -84,9 +84,9 @@ def get_location(self, v: Vertebra_Instance | int, poi: POI) -> tuple: subreg = Location.Additional_Vertebral_Body_Middle_Superior_Median if (v, subreg) in poi: return (v, subreg) - # Test if it has next POINT - if prev_vert is not None and (prev_vert, 50) in poi: - return (v, subreg, prev_vert, subreg) + # Fall back to averaging v's centroid with the previous vertebra's centroid. + if prev_vert is not None and (v, 50) in poi and (prev_vert, 50) in poi: + return (v, 50, prev_vert, 50) return (v, 50) def get_point(self, v: Vertebra_Instance | int, poi: POI) -> np.ndarray: @@ -395,10 +395,67 @@ def compute_lordosis_and_kyphosis(poi: POI, project_2D=True) -> dict[str, float return out +def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np.ndarray | None: + """Return an endplate-plane A/P direction from the buffered endplate POIs, or None. + + Uses the *relevant* endplate landmark for ``mv`` (``Vertebral_Body_Endplate_Superior`` + for :attr:`MoveTo.TOP`, ``Vertebral_Body_Endplate_Inferior`` for :attr:`MoveTo.BOTTOM`) + together with ``Vertebra_Direction_Right`` and ``Vertebra_Corpus`` to build a P/A + vector that lies in the specific endplate plane, rather than the averaged + vertebral body plane implied by ``Vertebra_Direction_Posterior``. + + The sign convention matches :func:`_get_norm` for ``Location.Vertebra_Direction_Posterior`` + with ``inv=1`` — the returned vector points anteriorly, so callers can multiply + by ``inv`` unchanged. + """ + if mv == MoveTo.TOP: + endplate_loc = Location.Vertebral_Body_Endplate_Superior + elif mv == MoveTo.BOTTOM: + endplate_loc = Location.Vertebral_Body_Endplate_Inferior + else: + return None + if (vert, 50) not in poi or (vert, endplate_loc) not in poi or (vert, Location.Vertebra_Direction_Right) not in poi: + return None + corpus = np.array(poi[vert, 50], dtype=float) + ep = np.array(poi[vert, endplate_loc], dtype=float) + r_pt = np.array(poi[vert, Location.Vertebra_Direction_Right], dtype=float) + n = ep - corpus + if mv == MoveTo.BOTTOM: + n = -n # flip inferior endplate so both cases point superior + n_norm = np.linalg.norm(n) + r_vec = r_pt - corpus + r_norm = np.linalg.norm(r_vec) + if n_norm < 1e-8 or r_norm < 1e-8: + return None + n /= n_norm + r_vec /= r_norm + # cross(right, superior-pointing) lies in the endplate plane and points posterior; + # negate to match _get_norm's default sign (anterior for inv=1). + p = -np.cross(r_vec, n) + p_norm = np.linalg.norm(p) + if p_norm < 1e-8: + return None + return p / p_norm + + def _get_norm(poi: POI, id1: int | Vertebra_Instance, mv: MoveTo, location: Location, inv: int = 1) -> np.ndarray | None: # noqa: ARG001 - """Return the normalised direction vector from a location POI to the vertebra centroid.""" + """Return the normalised direction vector from a location POI to the vertebra centroid. + + When ``location`` is :attr:`Location.Vertebra_Direction_Posterior` and ``mv`` targets + an endplate (:attr:`MoveTo.TOP` / :attr:`MoveTo.BOTTOM`), the buffered per-endplate + landmark (``Vertebral_Body_Endplate_Superior`` / ``_Inferior``) is preferred over + the averaged vertebral-body posterior direction — this yields the classical + endplate-line orientation used in Cobb-style lordosis/kyphosis measurements. + The endplate direction is used only when both the relevant endplate point and + ``Vertebra_Direction_Right`` are present in ``poi`` for that vertebra; otherwise + the code falls back to the WK-based averaged direction below. + """ if isinstance(id1, int): id1 = Vertebra_Instance(id1) + if location == Location.Vertebra_Direction_Posterior and mv in (MoveTo.TOP, MoveTo.BOTTOM): + ep_norm = _endplate_ap_direction(poi, id1, mv) + if ep_norm is not None: + return ep_norm * inv subreg = 50 if location in [Location.Vertebra_Disc_Inferior, Location.Vertebra_Disc_Superior]: subreg = 100 diff --git a/TPTBox/spine/spinestats/poi_fun/endplates.py b/TPTBox/spine/spinestats/poi_fun/endplates.py index fde98bba..b14ec2bd 100644 --- a/TPTBox/spine/spinestats/poi_fun/endplates.py +++ b/TPTBox/spine/spinestats/poi_fun/endplates.py @@ -399,10 +399,26 @@ def calc_endplate_points_( poi.info.setdefault("angle_superior_endplate", {}) poi.info.setdefault("angle_inferior_endplate", {}) + # Register direction-vector fields (auto-transformed by reorient / resample / + # to_cord_system) and additional per-label scalar fields (whose keys are + # remapped by map_labels). + from TPTBox.core.poi_fun.vector_fields import POI_INFO_LABEL_KEYED_FIELDS_KEY, POI_INFO_VECTOR_FIELDS_KEY + + _vec_fields = poi.info.setdefault(POI_INFO_VECTOR_FIELDS_KEY, []) + for _f in ("angle_superior_endplate", "angle_inferior_endplate"): + if _f not in _vec_fields: + _vec_fields.append(_f) + _lbl_fields = poi.info.setdefault(POI_INFO_LABEL_KEYED_FIELDS_KEY, []) + for _f in ("endplate_internal_angle",): + if _f not in _lbl_fields: + _lbl_fields.append(_f) if compute_curvature: poi.info.setdefault("curvature_superior_endplate", {}) poi.info.setdefault("curvature_inferior_endplate", {}) + for _f in ("curvature_superior_endplate", "curvature_inferior_endplate"): + if _f not in _lbl_fields: + _lbl_fields.append(_f) # Collect normals per vertebra so we can compute the inter-endplate # angle once both superior and inferior have been processed. diff --git a/docs/api/poi_fun.md b/docs/api/poi_fun.md index d1050ca2..b70a788f 100644 --- a/docs/api/poi_fun.md +++ b/docs/api/poi_fun.md @@ -44,3 +44,66 @@ from segmentation volumes. options: show_source: true filters: ["!^_"] + +## Direction-Vector Fields in `POI.info` + +`POI.info` can hold *auxiliary* per-vertebra data alongside the main POI points — +direction vectors (e.g. endplate PCA normals) and scalar metadata (e.g. per-vertebra +wedge angles or curvature). Producers register field names under two well-known +keys in `poi.info`, and `POI.reorient` / `POI.resample_from_to` / +`POI_Global.to_cord_system` / `POI.map_labels` then keep those fields aligned +with the POI points automatically. + +**Two registries:** + +- `poi.info["_vector_fields"]` (`POI_INFO_VECTOR_FIELDS_KEY`) — a list of field + names whose values are `{key: (x, y, z)}` dicts holding **unit direction + vectors in mm-space aligned with `poi.orientation`**. These fields are + auto-transformed by `reorient`, `resample_from_to`, and + `POI_Global.to_cord_system`, and their keys are also remapped by `map_labels`. + `rescale` is a no-op for them. +- `poi.info["_label_keyed_fields"]` (`POI_INFO_LABEL_KEYED_FIELDS_KEY`) — a + list of field names whose values are `{key: }` dicts (typically + scalars). Only their **keys** are remapped by `map_labels`; no + orientation-based transform is applied. + +Keys in either type of field may be integer region labels (`20`) or +`Vertebra_Instance`-name strings (`"L1"`); both forms are supported. + +**Producer pattern:** + +```python +from TPTBox.core.poi_fun.vector_fields import ( + POI_INFO_VECTOR_FIELDS_KEY, + POI_INFO_LABEL_KEYED_FIELDS_KEY, +) + +# Direction vectors (auto-rotated on reorient / resample): +vec_fields = poi.info.setdefault(POI_INFO_VECTOR_FIELDS_KEY, []) +if "my_vector_field" not in vec_fields: + vec_fields.append("my_vector_field") +poi.info["my_vector_field"] = {"L1": (0.09, -0.99, -0.02), ...} + +# Scalar per-vertebra metadata (key-remapped by map_labels only): +lbl_fields = poi.info.setdefault(POI_INFO_LABEL_KEYED_FIELDS_KEY, []) +if "my_scalar_field" not in lbl_fields: + lbl_fields.append("my_scalar_field") +poi.info["my_scalar_field"] = {"L1": 3.14, ...} +``` + +**Caveats:** + +- `resample_from_to` uses the real `R_ref.T @ R_self` rotation of the two + affines (safe for skewed grids); axcode-only heuristics are deliberately + avoided. +- `map_labels` handles duplicates with a *last-write-wins* policy after remap. +- `map_labels`' `label_map_full` (mapping `(region, subreg)` tuples) does NOT + trigger the key-remap of these fields — only `label_map_region` does. + +::: TPTBox.core.poi_fun.vector_fields + options: + show_source: true + filters: ["!^_"] + members: + - POI_INFO_VECTOR_FIELDS_KEY + - POI_INFO_LABEL_KEYED_FIELDS_KEY diff --git a/unit_tests/test_poi_vector_fields.py b/unit_tests/test_poi_vector_fields.py new file mode 100644 index 00000000..46cb90ca --- /dev/null +++ b/unit_tests/test_poi_vector_fields.py @@ -0,0 +1,200 @@ +"""Unit tests for auto-transform of direction-vector and label-keyed fields in ``POI.info``. + +Covers: +- ``POI.reorient`` — signed permutation of direction vectors. +- ``POI.resample_from_to`` — full rotation (``R_ref.T @ R_self``). +- ``POI.rescale`` — no-op for mm-space vectors. +- ``POI_Global.to_cord_system`` — RAS <-> LPS axis flips. +- ``POI.map_labels`` — key remap for both vector and label-keyed fields. +""" + +from __future__ import annotations + +import sys +import unittest +from pathlib import Path + +file = Path(__file__).resolve() +sys.path.append(str(file.parents[2])) + +import numpy as np # noqa: E402 + +from TPTBox.core.poi import POI # noqa: E402 +from TPTBox.core.poi_fun.vector_fields import ( # noqa: E402 + POI_INFO_LABEL_KEYED_FIELDS_KEY, + POI_INFO_VECTOR_FIELDS_KEY, +) + + +def _make_poi(orientation=("P", "I", "R"), zoom=(1.0, 1.0, 1.0)) -> POI: + """Build a small POI with identity rotation and one L1 corpus point.""" + centroids: dict[int, dict[int, tuple[float, float, float]]] = {20: {50: (1.0, 2.0, 3.0)}} + return POI( + centroids, + orientation=orientation, + zoom=zoom, + shape=(100, 100, 100), + origin=(0.0, 0.0, 0.0), + rotation=np.eye(3), + ) + + +def _register(poi: POI, vec_fields=(), lbl_fields=()) -> None: + if vec_fields: + poi.info.setdefault(POI_INFO_VECTOR_FIELDS_KEY, []).extend(vec_fields) + if lbl_fields: + poi.info.setdefault(POI_INFO_LABEL_KEYED_FIELDS_KEY, []).extend(lbl_fields) + + +class Test_Vector_Field_Reorient(unittest.TestCase): + def test_pir_to_las_signed_permutation(self): + # In PIR, +x=P, +y=I, +z=R. In LAS, +x=L=-R, +y=A=-P, +z=S=-I. + # v_PIR = (P=0.020, I=-0.092, R=0.996). Same physical direction in LAS is: + # L=-R=-0.996, A=-P=-0.020, S=-I=0.092 -> (-0.996, -0.020, 0.092). + poi = _make_poi(orientation=("P", "I", "R")) + _register(poi, vec_fields=["v"]) + v_in = np.array([0.02, -0.09, 0.996]) + v_in = v_in / np.linalg.norm(v_in) + poi.info["v"] = {"L1": tuple(v_in)} + out = poi.reorient(("L", "A", "S")) + got = np.asarray(out.info["v"]["L1"]) + expected = np.array([-v_in[2], -v_in[0], -v_in[1]]) + np.testing.assert_allclose(got, expected, atol=1e-9) + # Norm is preserved by a signed permutation. + self.assertAlmostEqual(float(np.linalg.norm(got)), 1.0, places=6) + # Full round-trip: PIR -> LAS -> PIR equals original. + back = out.reorient(("P", "I", "R")) + np.testing.assert_allclose(back.info["v"]["L1"], v_in, atol=1e-9) + + def test_unregistered_field_untouched(self): + poi = _make_poi() + # not registered + poi.info["orphan_vec"] = {"L1": (1.0, 0.0, 0.0)} + out = poi.reorient(("L", "A", "S")) + self.assertEqual(tuple(out.info["orphan_vec"]["L1"]), (1.0, 0.0, 0.0)) + + def test_none_or_malformed_values_skipped(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + poi.info["v"] = {"L1": (0.0, 1.0, 0.0), "L2": None, "L3": (1.0, 2.0)} + out = poi.reorient(("L", "A", "S")) + self.assertIsNone(out.info["v"]["L2"]) + # 2-tuple stays unchanged, since shape mismatch is silently skipped + self.assertEqual(tuple(out.info["v"]["L3"]), (1.0, 2.0)) + # 3-tuple got transformed + self.assertEqual(len(out.info["v"]["L1"]), 3) + + +class Test_Vector_Field_Rescale(unittest.TestCase): + def test_rescale_leaves_vectors_untouched(self): + poi = _make_poi(zoom=(0.5, 0.5, 3.0)) + _register(poi, vec_fields=["v"]) + v = (0.1, -0.9, 0.05) + poi.info["v"] = {"L1": v} + out = poi.rescale((1.0, 1.0, 1.0)) + np.testing.assert_allclose(out.info["v"]["L1"], v, atol=1e-12) + + +class Test_Vector_Field_ResampleFromTo(unittest.TestCase): + def test_identity_grid_is_noop(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + v = (0.09, -0.99, -0.02) + poi.info["v"] = {"L1": v} + # ref = a copy of the same POI -> same rotation -> R_ref.T @ R_self = I + ref = poi.copy() + out = poi.resample_from_to(ref) + np.testing.assert_allclose(out.info["v"]["L1"], v, atol=1e-6) + + def test_rotated_ref_applies_rotation(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + v = (1.0, 0.0, 0.0) + poi.info["v"] = {"L1": v} + # Target POI has a 90-degree z-rotation + R_z90 = np.array([[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]]) + ref = poi.copy() + ref.rotation = R_z90 + out = poi.resample_from_to(ref) + # v_target = R_ref.T @ R_self @ v = R_z90.T @ I @ [1,0,0] = R_z90.T @ [1,0,0] = [0, -1, 0] + np.testing.assert_allclose(out.info["v"]["L1"], (0.0, -1.0, 0.0), atol=1e-9) + + +class Test_ToCordSystem(unittest.TestCase): + def test_ras_to_lps_flips_xy_only(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + poi.info["v"] = {"L1": (0.1, 0.2, 0.3)} + g = poi.to_global() # RAS by default + self.assertFalse(g.itk_coords) + self.assertEqual(tuple(g.info["v"]["L1"]), (0.1, 0.2, 0.3)) + g_itk = g.to_cord_system(itk_coords=True) + np.testing.assert_allclose(g_itk.info["v"]["L1"], (-0.1, -0.2, 0.3), atol=1e-12) + # roundtrip + g_back = g_itk.to_cord_system(itk_coords=False) + np.testing.assert_allclose(g_back.info["v"]["L1"], (0.1, 0.2, 0.3), atol=1e-12) + + +class Test_MapLabels(unittest.TestCase): + def test_vector_field_keys_remapped(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + v = (0.1, -0.9, 0.4) + poi.info["v"] = {"L1": v} + out = poi.map_labels(label_map_region={20: 2}) # L1 -> C2 + self.assertNotIn("L1", out.info["v"]) + self.assertIn("C2", out.info["v"]) + np.testing.assert_allclose(out.info["v"]["C2"], v, atol=1e-12) + + def test_label_keyed_scalar_field_remapped(self): + poi = _make_poi() + _register(poi, lbl_fields=["endplate_internal_angle"]) + poi.info["endplate_internal_angle"] = {"L1": 3.14} + out = poi.map_labels(label_map_region={20: 2}) + self.assertNotIn("L1", out.info["endplate_internal_angle"]) + self.assertAlmostEqual(out.info["endplate_internal_angle"]["C2"], 3.14) + + def test_int_keys_also_remapped(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + poi.info["v"] = {20: (0.0, 1.0, 0.0), 21: (1.0, 0.0, 0.0)} + out = poi.map_labels(label_map_region={20: 2}) + self.assertNotIn(20, out.info["v"]) + self.assertIn(2, out.info["v"]) + # unaffected key + self.assertIn(21, out.info["v"]) + np.testing.assert_allclose(out.info["v"][2], (0.0, 1.0, 0.0), atol=1e-12) + + def test_unknown_string_key_kept(self): + poi = _make_poi() + _register(poi, lbl_fields=["s"]) + poi.info["s"] = {"NOT_A_VERT": 42.0, "L1": 1.0} + out = poi.map_labels(label_map_region={20: 2}) + self.assertIn("NOT_A_VERT", out.info["s"]) + self.assertNotIn("L1", out.info["s"]) + self.assertIn("C2", out.info["s"]) + + def test_empty_map_is_noop(self): + poi = _make_poi() + _register(poi, vec_fields=["v"]) + v = (0.1, -0.9, 0.4) + poi.info["v"] = {"L1": v} + out = poi.map_labels(label_map_region={}) + self.assertEqual(tuple(out.info["v"]["L1"]), v) + + +class Test_Composition(unittest.TestCase): + def test_reorient_and_map_labels_commute(self): + # For a vector-transform + key-remap, order shouldn't matter. + poi_a = _make_poi() + poi_b = _make_poi() + for poi in (poi_a, poi_b): + _register(poi, vec_fields=["v"]) + poi.info["v"] = {"L1": (0.02, -0.09, 0.996)} + out_a = poi_a.reorient(("L", "A", "S")).map_labels(label_map_region={20: 2}) + out_b = poi_b.map_labels(label_map_region={20: 2}).reorient(("L", "A", "S")) + np.testing.assert_allclose(out_a.info["v"]["C2"], out_b.info["v"]["C2"], atol=1e-9) + + +if __name__ == "__main__": + unittest.main() From 431a3203f7e5328c3830cfd20e2c5f2e53f34c74 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 14:57:17 +0000 Subject: [PATCH 17/26] claude_remapping --- TPTBox/core/poi_fun/poi_abstract.py | 20 +-- TPTBox/core/poi_fun/vector_fields.py | 131 +++++++++-------- TPTBox/core/poi_fun/vertebra_direction.py | 27 +++- TPTBox/spine/spinestats/_load_nako_wh.py | 11 +- TPTBox/spine/spinestats/_run_all.py | 13 +- TPTBox/spine/spinestats/veridah_angles.py | 102 +++++++++++-- docs/api/poi_fun.md | 63 ++++++-- unit_tests/test_poi_vector_fields.py | 167 ++++++++++++++++++++++ 8 files changed, 433 insertions(+), 101 deletions(-) diff --git a/TPTBox/core/poi_fun/poi_abstract.py b/TPTBox/core/poi_fun/poi_abstract.py index fbc8cff4..6139891f 100755 --- a/TPTBox/core/poi_fun/poi_abstract.py +++ b/TPTBox/core/poi_fun/poi_abstract.py @@ -636,23 +636,17 @@ def map_labels( continue poi_new[region:subreg] = value new_values = poi_new - from TPTBox.core.poi_fun.vector_fields import _remap_label_name_inplace, _remap_vector_field_keys_inplace + if new_values is None: + # No mapping happened (all input maps were empty) -> info stays as is. + return self if inplace else self.copy() - def _sync_info(target_info: dict) -> None: - _remap_vector_field_keys_inplace(target_info, label_map_region_) - _remap_label_name_inplace(target_info, label_map_region_, label_map_subregion_) + from TPTBox.core.poi_fun.vector_fields import _remap_vector_field_keys_inplace - if new_values is None: - out = self if inplace else self.copy() - _sync_info(out.info) - return out + target = self if inplace else self.copy(centroids=new_values) if inplace: self.centroids = new_values - _sync_info(self.info) - return self - out = self.copy(centroids=new_values) - _sync_info(out.info) - return out + _remap_vector_field_keys_inplace(target.info, label_map_region_, label_map_subregion_) + return target def map_labels_( self, diff --git a/TPTBox/core/poi_fun/vector_fields.py b/TPTBox/core/poi_fun/vector_fields.py index f21b9d07..078cc693 100644 --- a/TPTBox/core/poi_fun/vector_fields.py +++ b/TPTBox/core/poi_fun/vector_fields.py @@ -60,86 +60,99 @@ def _transform_direction_vectors_inplace(info: dict, trans: np.ndarray) -> None: vectors[k] = tuple(float(x) for x in new_v) -def _remap_vector_field_keys_inplace(info: dict, region_map: dict) -> None: - """Remap the top-level keys of every registered label-keyed field via ``region_map``. +def _map_one_key(k, key_map: dict | None): + """Map a single dict key via ``key_map``. + + ``int`` keys are looked up directly. ``str`` keys resolve through + ``Vertebra_Instance`` (name -> value) and, on hit, the mapped value is + turned back into the corresponding ``Vertebra_Instance`` name (or kept + as ``str(id)`` when the target isn't a known instance). + """ + if not key_map: + return k + if isinstance(k, int) and k in key_map: + return key_map[k] + if isinstance(k, str): + try: + label = Vertebra_Instance[k].value + except KeyError: + return k + if label in key_map: + try: + return Vertebra_Instance(key_map[label]).name + except ValueError: + return str(key_map[label]) + return k + + +def _remap_vector_field_keys_inplace(info: dict, region_map: dict | None, subregion_map: dict | None = None) -> None: + """Remap keys of every registered label-keyed field in ``info``. Considers fields registered under both :data:`POI_INFO_VECTOR_FIELDS_KEY` (direction vectors) and :data:`POI_INFO_LABEL_KEYED_FIELDS_KEY` (scalar - per-label fields like ``endplate_internal_angle`` or ``curvature_*``). - Keys may be either integer region labels or ``Vertebra_Instance``-name - strings ("L1", "T12", ...); both are matched against ``region_map`` - (int-keyed). Duplicates after remapping keep the last write. No-op if - no fields are registered or ``region_map`` is empty. + per-label fields), plus ``label_name`` (always handled; migrated to the + nested form on the fly). + + Behaviour is dispatched by *value type*: + + - **Flat fields** (values are tuples / scalars): only outer keys are + remapped via ``region_map``. + - **Nested fields** (values are dicts, e.g. ``label_name`` / + ``{region: {subregion: name, "name": group}}``): outer keys are + remapped via ``region_map``, inner keys via ``subregion_map``, and the + special ``"name"`` group entry is preserved. When two source regions + collide onto one target, their inner dicts merge (last-write-wins on + overlapping keys). + + Keys may be integer labels or ``Vertebra_Instance``-name strings; both + are matched against the int-keyed maps. No-op if there is nothing to do. """ - if not region_map: + from TPTBox.core.poi_fun.poi_abstract import LABEL_NAME, _GROUP_NAME_KEY, label_name_dict + + if not region_map and not subregion_map: return field_names: list[str] = [] for key in (POI_INFO_VECTOR_FIELDS_KEY, POI_INFO_LABEL_KEYED_FIELDS_KEY): names = info.get(key) if names: field_names.extend(names) + # label_name is always handled -- ensure the nested-form migration runs, then include it. + if LABEL_NAME in info and LABEL_NAME not in field_names: + label_name_dict(info) + field_names.append(LABEL_NAME) if not field_names: return for name in field_names: vectors = info.get(name) if not isinstance(vectors, dict): continue - remapped = {} + remapped: dict = {} for k, v in vectors.items(): - new_k = k - if isinstance(k, int) and k in region_map: - new_k = region_map[k] - elif isinstance(k, str): - try: - label = Vertebra_Instance[k].value - except KeyError: - label = None - if label is not None and label in region_map: - try: - new_k = Vertebra_Instance(region_map[label]).name - except ValueError: - new_k = k - remapped[new_k] = v + new_k = _map_one_key(k, region_map) + if new_k is None: + continue # drop entries whose region is mapped to None + new_v = v + # Nested field: recurse into inner dict. + if isinstance(v, dict): + new_inner: dict = {} + for ik, iv in v.items(): + if ik == _GROUP_NAME_KEY: + new_inner[_GROUP_NAME_KEY] = iv + continue + new_ik = _map_one_key(ik, subregion_map) + if new_ik is None: + continue # drop entries whose subregion is mapped to None + new_inner[new_ik] = iv + new_v = new_inner + # Merge on outer-key collision when both values are dicts (label_name-style). + if new_k in remapped and isinstance(remapped[new_k], dict) and isinstance(new_v, dict): + remapped[new_k].update(new_v) + else: + remapped[new_k] = new_v vectors.clear() vectors.update(remapped) -def _remap_label_name_inplace(info: dict, region_map: dict | None, subregion_map: dict | None) -> None: - """Remap the region + subregion keys of ``info["label_name"]`` in place. - - ``label_name`` uses the nested format - ``{region:int -> {subregion:int -> name:str, "name": group_name:str}}`` - (see :func:`normalize_label_name`). This helper remaps top-level region keys - via ``region_map`` and, for each inner dict, remaps subregion keys via - ``subregion_map``. The special ``"name"`` group-name entry is preserved. - No-op if the field is absent or both maps are empty. - """ - from TPTBox.core.poi_fun.poi_abstract import LABEL_NAME, _GROUP_NAME_KEY, label_name_dict - - if not region_map and not subregion_map: - return - if LABEL_NAME not in info: - return - ln = label_name_dict(info) # ensures nested form - remapped: dict[int, dict] = {} - for region, inner in ln.items(): - new_region = region_map[region] if region_map and region in region_map else region - new_inner: dict = {} - for k, v in inner.items(): - if k == _GROUP_NAME_KEY: - new_inner[_GROUP_NAME_KEY] = v - elif subregion_map and k in subregion_map: - new_inner[subregion_map[k]] = v - else: - new_inner[k] = v - # merge if two source regions collide onto one target (last-write-wins on inner keys). - if new_region in remapped: - remapped[new_region].update(new_inner) - else: - remapped[new_region] = new_inner - info[LABEL_NAME] = remapped - - def _rotate_direction_vectors_inplace(info: dict, src_rot, tgt_rot) -> None: """Rotate registered direction-vector fields from ``src_rot`` to ``tgt_rot``. diff --git a/TPTBox/core/poi_fun/vertebra_direction.py b/TPTBox/core/poi_fun/vertebra_direction.py index 7a510819..f2d3b043 100644 --- a/TPTBox/core/poi_fun/vertebra_direction.py +++ b/TPTBox/core/poi_fun/vertebra_direction.py @@ -147,15 +147,26 @@ def calc_orientation_of_vertebra_PIR( if last_vert == 20: last_vert = None break - max_vert_key = max(vert_keys) + # Only pre-sacral vertebrae influence the "next free label" search — S1 (and higher) + # may be present in poi_iso via `calc_endplate_points_` (Sacrum_Endplate landmark), + # but their label must NOT be counted here, otherwise spline anchors would land at + # labels 27/28 and get promoted to phantom vertebrae by the downstream + # `poi_iso.extract_subregion(source_subreg_point_id)` loop (which would then write + # spurious direction POIs at those labels and fool `_get_last_thoracic` into + # returning T13 / COCC). + presacral_keys = [k for k in vert_keys if k < Vertebra_Instance.S1.value] + max_vert_key = max(presacral_keys) if presacral_keys else max(vert_keys) + anchor_labels: set[int] = set() if (Vertebra_Instance.S1.value, Location.Vertebral_Body_Endplate_Superior.value) in poi_iso: poi_iso[max_vert_key + 1, spline_subreg_point_id] = poi_iso[ (Vertebra_Instance.S1.value, Location.Vertebral_Body_Endplate_Superior.value) ] max_vert_key += 1 + anchor_labels.add(max_vert_key) if last_vert is not None and (last_vert, Location.Vertebral_Body_Endplate_Inferior.value) in poi_iso: poi_iso[max_vert_key + 1, spline_subreg_point_id] = poi_iso[(last_vert, Location.Vertebral_Body_Endplate_Inferior.value)] max_vert_key += 1 + anchor_labels.add(max_vert_key) ##### # spline: body_spline, body_spline_der = poi_iso.fit_spline(location=spline_subreg_point_id, vertebra=True) @@ -184,6 +195,10 @@ def calc_orientation_of_vertebra_PIR( fill_back = out.copy() if do_fill_back else None # Draw a plain with the up_vector an cut it with intersection_target for reg_label, _, cords in poi_iso.extract_subregion(source_subreg_point_id).items(): + # Spline-anchor labels (added above to influence the spline fit) must not be + # processed here — they'd otherwise get written back to `ret` as pseudo-vertebrae. + if reg_label in anchor_labels: + continue # calculate_normal_vector if reg_label in down_vector: normal_vector_down = down_vector[reg_label] @@ -222,7 +237,15 @@ def calc_orientation_of_vertebra_PIR( arr = subreg_sar.set_array(fill_back).reorient(poi.orientation).rescale_(poi.zoom).get_array() fill_back_nii.set_array_(arr) - ret = calc_centroids(subreg_iso.set_array(out), second_stage=subreg_id, extend_to=poi_iso.copy(), inplace=True) + # Strip spline anchors from the working POI before extending: they were only + # needed to influence `fit_spline` above, and must not leak into the returned POI + # as pseudo-vertebrae (they'd otherwise show up as phantom regions labelled 25/26 + # / 27/28, fooling `_get_last_thoracic` etc. into treating them as real). + poi_iso_clean = poi_iso.copy() + for _a in anchor_labels: + if (_a, spline_subreg_point_id.value) in poi_iso_clean: + del poi_iso_clean.centroids[_a, spline_subreg_point_id.value] + ret = calc_centroids(subreg_iso.set_array(out), second_stage=subreg_id, extend_to=poi_iso_clean, inplace=True) poi._vert_orientation_pir = {} if save_normals_in_info: diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index 6cbb1a56..fa18c64d 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -154,6 +154,14 @@ def resolve_pick( # (``main-conflict:pd``, ``mevibe-conflict:pd``, ``vibe-conflict:pd``, …). idx = max(range(len(labels)), key=labels.__getitem__) return candidates[idx] # transient: do not save + if key == "main:T2haste" or key.endswith(":T2haste"): + # auto-accept: when duplicates differ only in the presence of a ``sequ`` entity + # (one raw file without sequ, one or more with an explicit sequ number for the same + # acquisition), prefer the sequ-numbered candidate. Skips the prompt only when this + # split is unambiguous (exactly one sequ'd candidate); otherwise falls through. + with_sequ = [c for c in candidates if getattr(c, "get", lambda *_: None)("sequ", None) is not None] + if len(with_sequ) == 1 and len(with_sequ) < len(candidates): + return with_sequ[0] # transient: do not save choice, reason = _prompt_choice(sub, key, question, labels, allow_discard=allow_discard) if choice == "__skip__": return candidates[0] if candidates else None # transient: do not save @@ -1096,6 +1104,7 @@ def hard_link( info = {"run": None} if isinstance(bf, str): bf = BIDS_FILE(bf, dataset) + bf.info.pop("run", None) assert len([k for k, v in bf.loop_keys() if k not in allowed_keys]) == 0, ( [k for k, v in bf.loop_keys() if k not in allowed_keys], bf, @@ -1564,7 +1573,7 @@ def _drain(fs): "Pass '' to disable.", ) args = parser.parse_args() - test = True + test = False if args.build_corrected_index: build_corrected_index() diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 75059029..8e7127b9 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -512,10 +512,21 @@ def _need(*keys: str, compute: bool) -> bool: if need_veridah and poi is not None: try: - from TPTBox.spine.spinestats.veridah_angles import compute_veridah_variants + from TPTBox.spine.spinestats.veridah_angles import compute_veridah_variants, plot_veridah_variants logger.on_debug("veridah variants") out["curv_veridah"] = compute_veridah_variants(poi, file_dict.get("veridah")) + if file_dict.get("veridah") is not None and file_dict.get("vert") is not None: + veridah_jpg_out = t2w_bf.get_changed_path( + "jpg", "snp", "derivatives_spine_inference_162_sacrumfix_subregionmeasures-v2", info={"seg": "cobb-veridah"} + ) + plot_veridah_variants( + veridah_jpg_out, + poi, + file_dict["t2w"], + file_dict["vert"], + file_dict.get("veridah"), + ) except Exception: logger.on_fail("veridah variants error caught") logger.print_error() diff --git a/TPTBox/spine/spinestats/veridah_angles.py b/TPTBox/spine/spinestats/veridah_angles.py index 2c40bb54..72e0fa26 100644 --- a/TPTBox/spine/spinestats/veridah_angles.py +++ b/TPTBox/spine/spinestats/veridah_angles.py @@ -13,6 +13,13 @@ Both variants live under the top-level JSON key ``curv_veridah`` (see ``_run_all.py``); the standard ``curv`` output is untouched. + +:func:`plot_veridah_variants` writes a matching snapshot JPG. It re-uses +the already-computed POI points (no re-segmentation, no re-POI-derivation) +by driving the label re-key through :meth:`POI.map_labels` — which also +carries the registered ``label_name`` and vector-field metadata along — +and applies the same integer mapping to the vertebra NIfTI via +:meth:`NII.map_labels` so the overlay matches the relabeled points. """ from __future__ import annotations @@ -20,9 +27,11 @@ import json from pathlib import Path +from TPTBox.core.nii_wrapper import NII, Image_Reference, to_nii from TPTBox.core.poi import POI, POI_Descriptor from TPTBox.core.vert_constants import Vertebra_Instance -from TPTBox.spine.spinestats.angles import compute_lordosis_and_kyphosis +from TPTBox.spine.snapshot2D.snapshot_modular import create_snapshot +from TPTBox.spine.spinestats.angles import compute_lordosis_and_kyphosis, plot_compute_lordosis_and_kyphosis def _load_veridah(path: Path) -> dict | None: @@ -39,18 +48,7 @@ def _load_veridah(path: Path) -> dict | None: def _relabel_poi(poi: POI, mapping: dict[int, int]) -> POI: - """Return a copy of ``poi`` whose region ids are re-keyed via ``mapping``. - - Regions absent from ``mapping`` are dropped. Non-vertebra regions - (i.e. those already outside the vertebra label range) fall through - with their original id — but in practice a POI produced by - :func:`calc_poi_from_subreg_vert` only carries vertebra regions. - """ - new_centroids = POI_Descriptor() - for region, subregion, coord in poi.centroids.items(): - if region in mapping: - new_centroids[(mapping[region], subregion)] = coord - return poi.copy(centroids=new_centroids) + return poi.map_labels(label_map_region=mapping) def _last_k_relabel(poi: POI, k: int) -> POI: @@ -87,7 +85,7 @@ def _compute_anomaly_variant(poi: POI, veridah_json_path: Path | None) -> dict[s fpath = veridah.get("fpath") if not isinstance(orig, list) or not isinstance(fpath, list) or len(orig) != len(fpath): return None - mapping = {int(o): int(f) for o, f in zip(orig, fpath)} + mapping = {int(o): int(f) for o, f in zip(orig, fpath) if int(o) != int(f)} return compute_lordosis_and_kyphosis(_relabel_poi(poi, mapping)) @@ -100,6 +98,82 @@ def _compute_k_variant(poi: POI, k: int) -> dict[str, float | None]: return {"lumbar_lordosis": full.get("lumbar_lordosis")} +def _veridah_region_map(veridah_json_path: Path) -> dict[int, int] | None: + """Return the ``orig_label -> fpath`` region mapping from a VERIDAH stat json.""" + veridah = _load_veridah(Path(veridah_json_path)) + if veridah is None: + return None + orig = veridah.get("orig_label") + fpath = veridah.get("fpath") + if not isinstance(orig, list) or not isinstance(fpath, list) or len(orig) != len(fpath): + return None + return {int(o): int(f) for o, f in zip(orig, fpath) if int(o) != int(f)} + + +def plot_veridah_variants( + jpg_path: str | Path | None, + poi: POI, + img: Image_Reference, + seg_vert: Image_Reference, + veridah_json_path: Path | str | None, + line_len: int = 100, + project_2D: bool = True, +) -> tuple[dict[str, float | None] | None, str | None]: + """Render the VERIDAH-corrected lordosis / kyphosis snapshot next to the standard one. + + Re-uses the *already computed* POI points -- the angle numbers are cheap + trig on those points, and no re-segmentation / re-``calc_poi_from_subreg_vert`` + happens. The re-key is driven through :meth:`POI.map_labels`, so any + registered ``label_name`` / vector-field metadata is carried along + automatically. The vertebra segmentation is remapped with the same + integer table via :meth:`NII.map_labels` so the overlay matches. + + Args: + jpg_path: Output path for the snapshot JPG. If ``None``, the image is not saved. + poi: The (already computed) POI object. + img: The reference image on which to plot. + seg_vert: Vertebra segmentation NIfTI, remapped alongside the POI. + veridah_json_path: Path to the VERIDAH ``_stat.json`` giving + ``{orig_label, fpath}`` lists. + line_len: Length of the direction lines drawn on the snapshot. + project_2D: Compute the angles as 2D sagittal projections. + + Returns: + ``(anomaly_angles, saved_path)`` -- ``anomaly_angles`` is the same + dict :func:`compute_lordosis_and_kyphosis` returns (or ``None`` + when the VERIDAH file is missing/malformed), ``saved_path`` is + the resolved output path (or ``None`` when nothing was written). + """ + if veridah_json_path is None: + return None, None + mapping = _veridah_region_map(Path(veridah_json_path)) + if mapping is None: + return None, None + + relabeled_poi = _relabel_poi(poi, mapping) + if not relabeled_poi.centroids: + return None, None + # Remap the vertebra segmentation with the same integer table so the overlay + # numbering matches the (already-remapped) POI. NII.map_labels leaves labels + # that are not listed in ``mapping`` unchanged. + seg_nii = seg_vert if isinstance(seg_vert, NII) else to_nii(seg_vert, seg=True) + seg_relabeled = seg_nii.map_labels(mapping, verbose=False) # type: ignore[arg-type] + + angles, _frame = plot_compute_lordosis_and_kyphosis( + None, + relabeled_poi, + img, + seg_relabeled, + line_len=line_len, + project_2D=project_2D, + ) + saved: str | None = None + if jpg_path is not None: + create_snapshot(jpg_path, [_frame]) + saved = str(jpg_path) + return angles, saved + + def compute_veridah_variants(poi: POI, veridah_json_path: Path | str | None) -> dict: """Return the ``curv_veridah`` block for one subject. diff --git a/docs/api/poi_fun.md b/docs/api/poi_fun.md index 124a0e93..b4cd34ae 100644 --- a/docs/api/poi_fun.md +++ b/docs/api/poi_fun.md @@ -91,20 +91,61 @@ if "my_scalar_field" not in lbl_fields: poi.info["my_scalar_field"] = {"L1": 3.14, ...} ``` -**Special-cased field — `info["label_name"]`:** +**Assigning names (`info["label_name"]`):** -Human-readable per-point / per-region names live in the nested structure -`{region:int -> {subregion:int -> name:str, "name": group_name:str}}` (see -`poi_abstract.LABEL_NAME` and `normalize_label_name`). `map_labels` remaps this -field automatically via its own helper `_remap_label_name_inplace`, which: +Human-readable per-point and per-region names live in +`poi.info["label_name"]` as `{region: {subregion: name, "name": group_name}}`. +Use the accessors on `Abstract_POI` (available on both `POI` and `POI_Global`) +instead of writing the dict directly: -- remaps the top-level `region` keys via `label_map_region`; -- remaps the *inner* `subregion` keys via `label_map_subregion`; -- preserves the special `"name"` group-name entry; -- merges inner dicts with *last-write-wins* on inner-key collisions when two - source regions map onto the same target. +```python +poi.set_label_name(region=2, subregion=10, name="FLCPC") # per-point label +poi.set_level_one_name(region=2, name="Femur") # region group name + +poi.label_name(2, 10) # -> "FLCPC" +poi.level_one_name(2) # -> "Femur" +``` + +`region` / `subregion` accept `int`, numeric string, or `Enum` members. A +custom name in `label_name` always takes priority over the auto-derived name +from `level_one_info` / `level_two_info`; if none is set, `.label_name(...)` +falls back to the enum name, and finally to the raw id as a string. A warning +is emitted when a custom name conflicts with the `level_two_info` enum name +for the same id. + +`map_labels` remaps this field automatically: the same +`_remap_vector_field_keys_inplace` helper handles both flat fields and the +nested `label_name` structure by dispatching on value type. For `label_name`, +`label_map_region` remaps the top-level region keys, `label_map_subregion` +remaps the inner subregion keys, and the `"name"` group entry is preserved. +Two source regions colliding onto one target merge inner dicts with +*last-write-wins* on overlapping keys. No explicit registration required — +`label_name` is always handled. + +**Names in 3D Slicer:** + +`POI_Global.save_mrk(...)` writes a `.mrk.json` markup file whose control +points and groups inherit these names directly: + +```python +poi_global = poi.to_global() +poi_global.save_mrk("points.mrk.json", pointLabelsVisibility=True) +``` + +Per control point, `save_mkr.get_desc(poi, region, subregion)` looks up: + +- `label` — from `poi.info["label_name"][region][subregion]`; falls back to + the `level_two_info` enum name (or the raw subregion id). +- `name2` (group label shown in the markup tree) — from + `poi.info["label_name"][region]["name"]`; falls back to + `poi.info["label_group_name"][region]` and finally to the + `level_one_info` enum name. -No explicit registration is required for `label_name` — it is always handled. +So `poi.set_label_name(...)` and `poi.set_level_one_name(...)` are all you +need: Slicer displays those strings on hover, in the markup tree, and (when +`pointLabelsVisibility=True`) as 3D annotations. Enable +`split_by_region=True` on `save_mrk` to get one Slicer group per region +(named via `level_one_name`). **Caveats:** diff --git a/unit_tests/test_poi_vector_fields.py b/unit_tests/test_poi_vector_fields.py index e120752c..cac5fdb6 100644 --- a/unit_tests/test_poi_vector_fields.py +++ b/unit_tests/test_poi_vector_fields.py @@ -17,6 +17,8 @@ file = Path(__file__).resolve() sys.path.append(str(file.parents[2])) +from typing import ClassVar # noqa: E402 + import numpy as np # noqa: E402 from TPTBox.core.poi import POI # noqa: E402 @@ -258,6 +260,171 @@ def test_flat_legacy_format_migrated_and_remapped(self): self.assertEqual(ln[2][50], "L1_corpus") +class Test_LabelName_LegacyMigration(unittest.TestCase): + """The old flat ``{"(region, subreg)": name}`` format must migrate cleanly. + + Real-world fixture: a leg-atlas POI with 38 flat entries across 5 regions, + including multi-digit subregion ids like ``(2, 10)`` — those must survive + ``ast.literal_eval`` parsing. + """ + + _ATLAS_FLAT: ClassVar[dict[str, str]] = { + "(1, 1)": "TGT", "(1, 2)": "FHC", "(1, 3)": "FNC", "(1, 4)": "FAAP", + "(2, 1)": "FLCD", "(2, 2)": "FMCD", "(2, 3)": "FLCP", "(2, 4)": "FMCP", + "(2, 5)": "FNP", "(2, 6)": "FADP", "(2, 7)": "TGPP", "(2, 8)": "TGCP", + "(2, 9)": "FMCPC", "(2, 10)": "FLCPC", "(2, 11)": "TRMP", "(2, 12)": "TRLP", + "(3, 1)": "TLCL", "(3, 2)": "TMCM", "(3, 3)": "TKC", "(3, 4)": "TLCA", + "(3, 5)": "TLCP", "(3, 6)": "TMCA", "(3, 7)": "TMCP", "(3, 8)": "TTP", + "(3, 9)": "TAAP", "(3, 10)": "TMIT", "(3, 11)": "TLIT", + "(4, 1)": "FLM", "(4, 2)": "TMM", "(4, 3)": "TAC", "(4, 4)": "TADP", + "(5, 1)": "PPP", "(5, 2)": "PDP", "(5, 3)": "PMP", "(5, 4)": "PLP", + "(5, 5)": "PRPP", "(5, 6)": "PRDP", "(5, 7)": "PRHP", + } + + def test_normalize_label_name_migrates_flat_atlas(self): + from TPTBox.core.poi_fun.poi_abstract import normalize_label_name + + nested = normalize_label_name(dict(self._ATLAS_FLAT)) + # region keys are ints + self.assertEqual(set(nested.keys()), {1, 2, 3, 4, 5}) + # multi-digit inner keys survive parsing + self.assertEqual(nested[2][10], "FLCPC") + self.assertEqual(nested[2][12], "TRLP") + self.assertEqual(nested[3][11], "TLIT") + # inner keys are ints too, no leftover string keys + for region, inner in nested.items(): + for k in inner: + self.assertIsInstance(k, int, f"inner key {k!r} in region {region} is not int") + # count is preserved + total = sum(len(v) for v in nested.values()) + self.assertEqual(total, len(self._ATLAS_FLAT)) + + def test_normalize_is_idempotent(self): + from TPTBox.core.poi_fun.poi_abstract import normalize_label_name + + once = normalize_label_name(dict(self._ATLAS_FLAT)) + twice = normalize_label_name({**once}) + self.assertEqual(once, twice) + + def test_normalize_empty_and_none(self): + from TPTBox.core.poi_fun.poi_abstract import normalize_label_name + + self.assertEqual(normalize_label_name(None), {}) + self.assertEqual(normalize_label_name({}), {}) + + def test_label_name_dict_caches_migration_in_info(self): + from TPTBox.core.poi_fun.poi_abstract import label_name_dict + + info = {"label_name": dict(self._ATLAS_FLAT)} + out = label_name_dict(info) + # returned dict is the normalized form + self.assertEqual(out[1][1], "TGT") + # and the migrated form is cached back into info + self.assertIs(info["label_name"], out) + # a second call is a no-op (still nested) + out2 = label_name_dict(info) + self.assertIs(out2, info["label_name"]) + + def test_load_poi_migrates_atlas_from_disk(self): + """End-to-end: write the legacy JSON, load it, expect the nested form.""" + import json + import tempfile + + from TPTBox.core.poi import POI + + payload = [ + { + "direction": ["R", "A", "S"], + "zoom": [1.0, 1.0, 1.0], + "origin": [0.0, 0.0, 0.0], + "shape": [10, 10, 10], + "rotation": [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], + "format": "POI", + "label_name": dict(self._ATLAS_FLAT), + }, + # one dummy point so the file loads as a POI: {region: {subregion: (x,y,z)}} + {"1": {"1": [1.0, 2.0, 3.0]}}, + ] + with tempfile.NamedTemporaryFile(mode="w", suffix="_poi.json", delete=False) as f: + json.dump(payload, f) + path = f.name + try: + poi = POI.load(path) + ln = poi.info["label_name"] + self.assertEqual(ln[1][1], "TGT") + self.assertEqual(ln[2][10], "FLCPC") + # no more flat "(...)"-style keys + self.assertFalse(any(isinstance(k, str) and k.startswith("(") for k in ln)) + finally: + Path(path).unlink(missing_ok=True) + + def test_migration_then_map_labels(self): + """A flat-format POI still remaps correctly through map_labels.""" + poi = _make_poi() + poi.info["label_name"] = dict(self._ATLAS_FLAT) # legacy flat + # rename region 2 -> 20 and region 5 -> 50 in one shot + out = poi.map_labels(label_map_region={2: 20, 5: 50}) + ln = out.info["label_name"] + self.assertIn(20, ln) + self.assertIn(50, ln) + self.assertNotIn(2, ln) + self.assertNotIn(5, ln) + self.assertEqual(ln[20][10], "FLCPC") + self.assertEqual(ln[50][7], "PRHP") + + +class Test_LabelName_Accessors(unittest.TestCase): + """`set_label_name` / `set_level_one_name` write into ``info['label_name']`` + and get read back by ``label_name`` / ``level_one_name`` and by the + Slicer/mkr exporter.""" + + def _poi_with_enums(self) -> POI: + from TPTBox.core.vert_constants import Location, Vertebra_Instance + + poi = _make_poi() + poi.level_one_info = Vertebra_Instance + poi.level_two_info = Location + return poi + + def test_set_and_get_label_name(self): + poi = self._poi_with_enums() + poi.set_label_name(region=20, subregion=50, name="L1_corpus") + self.assertEqual(poi.label_name(20, 50), "L1_corpus") + # unset points fall back to the level_two_info enum name (or the raw id when + # no enum entry matches). + fallback = poi.label_name(21, 50) + self.assertIsInstance(fallback, str) + self.assertNotEqual(fallback, "L1_corpus") + + def test_set_and_get_level_one_name(self): + poi = self._poi_with_enums() + poi.set_level_one_name(region=20, name="Vertebra L1 custom") + self.assertEqual(poi.level_one_name(20), "Vertebra L1 custom") + # unset region falls back to the level_one_info enum name (L1 -> "L1") + self.assertEqual(poi.level_one_name(21), "L2") + + def test_set_label_name_persists_in_info(self): + poi = _make_poi() + poi.set_label_name(20, 50, "L1_corpus") + poi.set_level_one_name(20, "Femur") + ln = poi.info["label_name"] + self.assertEqual(ln[20][50], "L1_corpus") + self.assertEqual(ln[20]["name"], "Femur") + + def test_names_flow_into_slicer_export(self): + """`get_desc` (used by save_mrk) reads label/group name from label_name.""" + from TPTBox.core.poi_fun.save_mkr import get_desc + + poi = _make_poi() + poi.set_label_name(20, 50, "L1_corpus") + poi.set_level_one_name(20, "Spine") + g = poi.to_global() + name, name2, label = get_desc(g, region=20, subregion=50) + # `label` is the per-point custom name; `name2` is the region group name. + self.assertEqual(label, "L1_corpus") + self.assertEqual(name2, "Spine") + + class Test_Composition(unittest.TestCase): def test_reorient_and_map_labels_commute(self): # For a vector-transform + key-remap, order shouldn't matter. From 7d5b1a5b0a5350f4342a7f01ee700d70a2e9e5a1 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 16:17:17 +0000 Subject: [PATCH 18/26] fix angls. --- TPTBox/core/poi.py | 3 + TPTBox/spine/spinestats/_run_all.py | 108 ++++++++++++++++++---- TPTBox/spine/spinestats/angles.py | 83 +++++++++++++---- TPTBox/spine/spinestats/curvature.py | 2 +- TPTBox/spine/spinestats/veridah_angles.py | 20 ++-- 5 files changed, 170 insertions(+), 46 deletions(-) diff --git a/TPTBox/core/poi.py b/TPTBox/core/poi.py index 74c3bd75..470586a9 100755 --- a/TPTBox/core/poi.py +++ b/TPTBox/core/poi.py @@ -1401,6 +1401,9 @@ def calc_centroids( ctd_list[first_stage, int(i)] = out_coord else: ctd_list[int(i), second_stage] = out_coord + if extend_to is not None: + args.setdefault("level_one_info", extend_to.level_one_info) + args.setdefault("level_two_info", extend_to.level_two_info) return POI(ctd_list, **msk_nii._extract_affine(), **args) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 8e7127b9..c03a1142 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -23,8 +23,7 @@ from tqdm import tqdm -from TPTBox import Print_Logger -from TPTBox.core.bids_files import BIDS_FILE +from TPTBox import BIDS_FILE, POI, Print_Logger from TPTBox.core.dicom.dicom2nii_utils import load_json from TPTBox.core.internal.nii_help import save_json from TPTBox.core.nii_wrapper import to_nii @@ -32,6 +31,16 @@ DATASET_ROOT = Path("/DATA/NAS/datasets_processed/NAKO/dataset-nako") logger = Print_Logger() +# Version stamp written into every aggregate `_stat.json` at +# ``_provenance.version`` (built by :func:`_build_provenance`). Bump this whenever +# the produced numbers change in a way that requires already-cached files to be +# recomputed. Current bumps: +# 1 -- initial release (WK-direction based cobb/lordosis; no S1 sacrum endplate). +# 2 -- endplate-plane based lordosis / kyphosis (average of the two flanking +# endplates at each disc), S1 sacrum-endplate landmarks, and correct +# cranio-caudal ordering for the T13 annotation label. +CURRENT_VERSION = 2 + # Top-level keys we require inside a finished json before we consider a # subject "done" and skip recomputation. cobb/curv are optional and only # added when run_all is called with cobb=True. @@ -99,8 +108,7 @@ def get_nako_paths(nako_id: str) -> dict[str, Path | None]: veridah = None for suffix in ("VERIDAH-label-V2", "VERIDAH-label"): p = ( - DATASET_ROOT - / f"derivatives_spine_inference_162_sacrumfix/{pfx}/{sub}/T2w/" + DATASET_ROOT / f"derivatives_spine_inference_162_sacrumfix/{pfx}/{sub}/T2w/" f"sub-{sub}_sequ-stitched_acq-sag_mod-T2w_seg-vert_desc-{suffix}_stat.json" ) if p.exists(): @@ -137,6 +145,33 @@ def _segmentation_inputs(file_dict: dict) -> list[Path]: ] +def _stat_version(loaded_stat: dict | None) -> int: + """Return the ``_provenance.version`` of a loaded stat dict (defaults to 1).""" + if not isinstance(loaded_stat, dict): + return 1 + prov = loaded_stat.get("_provenance") + if isinstance(prov, dict): + try: + return int(prov.get("version", 1)) + except (TypeError, ValueError): + return 1 + return 1 + + +def _poi_is_stale_wrt_stat(stat_path: Path, loaded_stat: dict) -> bool: + """Return True when the POI buffer sitting next to ``stat_path`` should be rebuilt. + + Staleness is derived from the stat json's ``_provenance.version``: whenever + that reads < :data:`ANGLES_VERSION`, the sibling POI buffer is treated as + outdated (v1 stat + v1 POI travelled together). ``loaded_stat`` is the + already-loaded stat dict — pass ``{}`` if none exists (then nothing is stale + since there's no v1 marker to invalidate against). + """ + if not stat_path.exists() or not loaded_stat: + return False + return _stat_version(loaded_stat) < CURRENT_VERSION + + def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple[str, ...]) -> tuple[bool, dict | None]: """Return (valid, loaded_dict). @@ -164,6 +199,8 @@ def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple for k in required_keys: if k not in data: return False, None + if _stat_version(data) < CURRENT_VERSION: + return False, None return True, data @@ -249,7 +286,7 @@ def _build_provenance(file_dict: dict, poi_out: Path | str | None, prior: dict | if p.exists(): inputs["poi"] = _file_provenance(p, prior_inputs.get("poi")) return { - "version": 1, + "version": CURRENT_VERSION, "written_at": datetime.now(timezone.utc).isoformat(timespec="seconds"), "inputs": inputs, } @@ -349,6 +386,26 @@ def run_all( out = loaded except Exception: out = {} + # Version stale -> drop the ANGLE-related keys so the corresponding _need() + # checks trigger a recompute of just the affected metrics + JPGs, without + # invalidating the expensive body-composition / VBQ / muscle_fat / torso + # blocks that don't depend on the endplate-based lordosis fix. + _cur_ver = _stat_version(out) + if _cur_ver < 2: + logger.on_warning("version bump", _cur_ver, "->", CURRENT_VERSION, ": recomputing angle keys") + for _k in ( + "cobb", + "curv", + "curv_veridah", + "sva", + "coronal_balance", + "axial_rotation", + "segmental_endplate_angles", + "curvature_profile", + "multi_cobb", + "endplate_internal_angle", + ): + out.pop(_k, None) def _need(*keys: str, compute: bool) -> bool: return override or (any(k not in out for k in keys) and compute) @@ -420,7 +477,16 @@ def _need(*keys: str, compute: bool) -> bool: need_vibe_wf = need_mfi if not ( - need_cobb or need_ivd or need_vert or need_vbq or need_bcs or need_mfi or need_torso or need_curvature or need_pelvic or need_veridah + need_cobb + or need_ivd + or need_vert + or need_vbq + or need_bcs + or need_mfi + or need_torso + or need_curvature + or need_pelvic + or need_veridah ): if _merge_endplate_angles(out, Path(poi_out)) or save: out["_provenance"] = _build_provenance(file_dict, poi_out, out.get("_provenance")) @@ -440,13 +506,18 @@ def _need(*keys: str, compute: bool) -> bool: poi = None if need_poi: logger.on_debug("calc_poi_from_subreg_vert") - poi = calc_poi_from_subreg_vert( - vert, - spine, - subreg_id=[Location.Vertebra_Corpus, Location.Vertebra_Direction_Posterior, Location.Endplate, Location.Vertebra_Disc], - buffer_file=poi_out, - save_buffer_file=True, - ) + if _poi_is_stale_wrt_stat(final_out, out): + Path(poi_out).unlink(missing_ok=True) + if poi_out.exists(): + poi = POI.load(poi_out) + else: + poi = calc_poi_from_subreg_vert( + vert, + spine, + subreg_id=[Location.Vertebra_Corpus, Location.Vertebra_Direction_Posterior, Location.Endplate, Location.Vertebra_Disc], + buffer_file=poi_out, + save_buffer_file=True, + ) if need_cobb: try: project_2D = False @@ -579,6 +650,7 @@ def _need(*keys: str, compute: bool) -> bool: _merge_endplate_angles(out, Path(poi_out)) out["_provenance"] = _build_provenance(file_dict, poi_out, out.get("_provenance")) + out.pop("_version", None) # TODO can beremoved logger.on_save("save", final_out.name) save_json(final_out, out) return out @@ -975,7 +1047,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: OUT_FOLDER.mkdir(parents=True, exist_ok=True) N_CPUS = 1 # set >1 to parallelize OVERRIDE = False - aggregate = True + aggregate = False do_not_update = False test = False collector: ExcelCollector | None = None @@ -993,10 +1065,10 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: subjects = loop_over_repaired_nako(test=False, sort=aggregate) else: # subjects = tqdm(loop_over_repaired_nako(test=False, sort=aggregate), total=30645) - l = loop_over_repaired_nako(test=False, sort=aggregate) - # total = 1000 - # subjects = iter([next(l) for _ in range(total)]) - subjects = l + l = loop_over_repaired_nako(test=False, sort=True) # aggregate + total = 15 + subjects = iter([next(l) for _ in range(total)]) + # subjects = l # print(f"Run on {total=} random subset") if N_CPUS <= 1: diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index f5e41c66..4d724ba4 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -258,7 +258,7 @@ def compute_angel_between_two_points_( if id1 > id2: id1, id2 = id2, id1 # Reorient and rescale the POI data - poi.reorient_().rescale_() + poi.reorient_().rescale_(verbose=False) recompute_use_ivd_direction = False # Determine direction-specific settings location2 = None @@ -395,22 +395,17 @@ def compute_lordosis_and_kyphosis(poi: POI, project_2D=True) -> dict[str, float return out -def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np.ndarray | None: - """Return an endplate-plane A/P direction from the buffered endplate POIs, or None. - - Uses the *relevant* endplate landmark for ``mv`` (``Vertebral_Body_Endplate_Superior`` - for :attr:`MoveTo.TOP`, ``Vertebral_Body_Endplate_Inferior`` for :attr:`MoveTo.BOTTOM`) - together with ``Vertebra_Direction_Right`` and ``Vertebra_Corpus`` to build a P/A - vector that lies in the specific endplate plane, rather than the averaged - vertebral body plane implied by ``Vertebra_Direction_Posterior``. +def _single_endplate_ap_direction(poi: POI, vert: Vertebra_Instance, side: MoveTo) -> np.ndarray | None: + """A/P direction of *one* endplate (superior for TOP, inferior for BOTTOM). - The sign convention matches :func:`_get_norm` for ``Location.Vertebra_Direction_Posterior`` - with ``inv=1`` — the returned vector points anteriorly, so callers can multiply - by ``inv`` unchanged. + Uses ``Vertebra_Corpus``, ``Vertebral_Body_Endplate_Superior/Inferior`` and + ``Vertebra_Direction_Right`` to build a unit vector that lies in the endplate + plane and points *anteriorly* (matches :func:`_get_norm`'s ``inv=1`` sign + convention). Returns ``None`` if any required POI landmark is missing. """ - if mv == MoveTo.TOP: + if side == MoveTo.TOP: endplate_loc = Location.Vertebral_Body_Endplate_Superior - elif mv == MoveTo.BOTTOM: + elif side == MoveTo.BOTTOM: endplate_loc = Location.Vertebral_Body_Endplate_Inferior else: return None @@ -420,7 +415,7 @@ def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np. ep = np.array(poi[vert, endplate_loc], dtype=float) r_pt = np.array(poi[vert, Location.Vertebra_Direction_Right], dtype=float) n = ep - corpus - if mv == MoveTo.BOTTOM: + if side == MoveTo.BOTTOM: n = -n # flip inferior endplate so both cases point superior n_norm = np.linalg.norm(n) r_vec = r_pt - corpus @@ -438,6 +433,49 @@ def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np. return p / p_norm +def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np.ndarray | None: + """Return the endplate-plane A/P direction at a vertebra-disc *boundary*. + + For a lordosis / kyphosis chain to close at every transition (e.g. so + ``thoracic_kyphosis + lumbar_lordosis`` measures the same T4-top→L5-inferior + end-to-end angle as the total), the "bottom of upper" and "top of lower" at + each disc must use the *same* reference direction. This function averages the + two flanking endplate P/A directions at the disc: + + - :attr:`MoveTo.BOTTOM` at vertebra ``V`` averages ``V``'s inferior endplate + with the superior endplate of the next vertebra in the POI (``V.get_next_poi``). + - :attr:`MoveTo.TOP` at vertebra ``V`` averages ``V``'s superior endplate + with the inferior endplate of the previous vertebra (``V.get_previous_poi``). + + Falls back to the single-endplate direction when the neighbour is missing. + Returns ``None`` when even the local endplate is unavailable. + """ + if isinstance(vert, int): + vert = Vertebra_Instance(vert) + if mv == MoveTo.BOTTOM: + neighbour = vert.get_next_poi(poi) + neighbour_side = MoveTo.TOP + elif mv == MoveTo.TOP: + neighbour = vert.get_previous_poi(poi) + neighbour_side = MoveTo.BOTTOM + else: + return _single_endplate_ap_direction(poi, vert, mv) + own = _single_endplate_ap_direction(poi, vert, mv) + other = _single_endplate_ap_direction(poi, neighbour, neighbour_side) if neighbour is not None else None + # If the local endplate landmark is missing (calc_endplate_points_ ray-cast + # sometimes fails to hit the mask), mirror across the disc: use the neighbour's + # endplate as the direction proxy so both sides of the boundary agree. + if own is None: + return other + if other is None: + return own + avg = own + other + n = np.linalg.norm(avg) + if n < 1e-8: + return own + return avg / n + + def _get_norm(poi: POI, id1: int | Vertebra_Instance, mv: MoveTo, location: Location, inv: int = 1) -> np.ndarray | None: # noqa: ARG001 """Return the normalised direction vector from a location POI to the vertebra centroid. @@ -788,7 +826,7 @@ def plot_compute_lordosis_and_kyphosis( >>> print(angles) {'cervical_lordosis': 34.5, 'thoracic_kyphosis': 42.7, 'lumbar_lordosis': 50.3} """ - poi = poi.reorient().rescale_() + poi = poi.reorient().rescale_(verbose=False) poi = _add_artificial_ivd(poi) out = [] text_out = [] @@ -813,7 +851,16 @@ def plot_compute_lordosis_and_kyphosis( id1 = curvature_definition[name].get_start_vert(poi) id2 = curvature_definition[name].get_stop_vert(poi) - vert = round((id1.value + id2.value) / 2) + # Cranio-caudal midpoint via the anatomical order — arithmetic mean of + # `.value` breaks for T13 (value 28, ordered after T12 but numbered after + # S1/COCC), landing the annotation on L1 instead of somewhere thoracic. + order = Vertebra_Instance.order() + try: + i1, i2 = order.index(id1), order.index(id2) + mid_inst = order[(min(i1, i2) + max(i1, i2)) // 2] + vert = mid_inst.value + except ValueError: + vert = round((id1.value + id2.value) / 2) while (vert, 50) not in poi and vert != 0: vert -= 1 text_out.append((vert, (f"{v:.1f}° - {str(name).split('_')[-1]}", 25))) @@ -870,7 +917,7 @@ def plot_cobb_angle( >>> plot_cobb_angle("output.png", poi, img, seg, line_len=100, threshold_deg=10) """ - poi = poi.reorient().rescale_() + poi = poi.reorient().rescale_(verbose=False) poi = _add_artificial_ivd(poi) out = [] diff --git a/TPTBox/spine/spinestats/curvature.py b/TPTBox/spine/spinestats/curvature.py index 8f8f267f..3e63ff98 100644 --- a/TPTBox/spine/spinestats/curvature.py +++ b/TPTBox/spine/spinestats/curvature.py @@ -70,7 +70,7 @@ def _prep(poi: POI) -> POI: """Return the POI in the internal ``(P, I, R)`` orientation at 1 mm scale.""" - return poi.reorient().rescale() + return poi.reorient().rescale(verbose=False) def _get(poi: POI, vert: int | Vertebra_Instance, loc: Location) -> np.ndarray | None: diff --git a/TPTBox/spine/spinestats/veridah_angles.py b/TPTBox/spine/spinestats/veridah_angles.py index 72e0fa26..a607319c 100644 --- a/TPTBox/spine/spinestats/veridah_angles.py +++ b/TPTBox/spine/spinestats/veridah_angles.py @@ -71,10 +71,12 @@ def _last_k_relabel(poi: POI, k: int) -> POI: return _relabel_poi(poi, mapping) -def _compute_anomaly_variant(poi: POI, veridah_json_path: Path | None) -> dict[str, float | None] | None: +def _compute_anomaly_variant(poi: POI, veridah_json_path: Path | None, project_2D: bool = False) -> dict[str, float | None] | None: """Recompute the three regional angles after applying the VERIDAH label correction. - Returns ``None`` when the VERIDAH file is missing or unusable. + ``project_2D`` must match the setting used for the standard ``curv`` block + in :mod:`_run_all` (default ``False``, i.e. 3D angles). Returns ``None`` + when the VERIDAH file is missing or unusable. """ if veridah_json_path is None: return None @@ -86,15 +88,15 @@ def _compute_anomaly_variant(poi: POI, veridah_json_path: Path | None) -> dict[s if not isinstance(orig, list) or not isinstance(fpath, list) or len(orig) != len(fpath): return None mapping = {int(o): int(f) for o, f in zip(orig, fpath) if int(o) != int(f)} - return compute_lordosis_and_kyphosis(_relabel_poi(poi, mapping)) + return compute_lordosis_and_kyphosis(_relabel_poi(poi, mapping), project_2D=project_2D) -def _compute_k_variant(poi: POI, k: int) -> dict[str, float | None]: +def _compute_k_variant(poi: POI, k: int, project_2D: bool = False) -> dict[str, float | None]: """Compute lumbar-lordosis only, with the last ``k`` vertebrae treated as L1..Lk.""" relabeled = _last_k_relabel(poi, k) if not relabeled.centroids: return {"lumbar_lordosis": None} - full = compute_lordosis_and_kyphosis(relabeled) + full = compute_lordosis_and_kyphosis(relabeled, project_2D=project_2D) return {"lumbar_lordosis": full.get("lumbar_lordosis")} @@ -117,7 +119,7 @@ def plot_veridah_variants( seg_vert: Image_Reference, veridah_json_path: Path | str | None, line_len: int = 100, - project_2D: bool = True, + project_2D: bool = False, ) -> tuple[dict[str, float | None] | None, str | None]: """Render the VERIDAH-corrected lordosis / kyphosis snapshot next to the standard one. @@ -174,7 +176,7 @@ def plot_veridah_variants( return angles, saved -def compute_veridah_variants(poi: POI, veridah_json_path: Path | str | None) -> dict: +def compute_veridah_variants(poi: POI, veridah_json_path: Path | str | None, project_2D: bool = False) -> dict: """Return the ``curv_veridah`` block for one subject. Shape:: @@ -186,7 +188,7 @@ def compute_veridah_variants(poi: POI, veridah_json_path: Path | str | None) -> } """ veridah_path = Path(veridah_json_path) if veridah_json_path is not None else None - out: dict = {"anomaly": _compute_anomaly_variant(poi, veridah_path)} + out: dict = {"anomaly": _compute_anomaly_variant(poi, veridah_path, project_2D=project_2D)} for k in (4, 5, 6, 7): - out[f"k{k}"] = _compute_k_variant(poi, k) + out[f"k{k}"] = _compute_k_variant(poi, k, project_2D=project_2D) return out From cfa270aae09c73ebde34357df073d3dfaca17ac7 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Mon, 21 Sep 2026 16:25:15 +0000 Subject: [PATCH 19/26] logging --- TPTBox/spine/spinestats/_run_all.py | 6 +++--- TPTBox/spine/spinestats/angles.py | 4 ++-- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index c03a1142..0e2598f9 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -1066,9 +1066,9 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: else: # subjects = tqdm(loop_over_repaired_nako(test=False, sort=aggregate), total=30645) l = loop_over_repaired_nako(test=False, sort=True) # aggregate - total = 15 - subjects = iter([next(l) for _ in range(total)]) - # subjects = l + # total = 15 + # subjects = iter([next(l) for _ in range(total)]) + subjects = l # print(f"Run on {total=} random subset") if N_CPUS <= 1: diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 4d724ba4..5d703518 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -961,8 +961,8 @@ def plot_cobb_angle( if width < 50: padd = [(0, 0) for _ in range(3)] padd[axis] = (int(50 - width), int(50 - width)) - img = to_nii(img).apply_pad(padd) - seg = to_nii(seg, True).apply_pad(padd) + img = to_nii(img).apply_pad(padd, verbose=False) + seg = to_nii(seg, True).apply_pad(padd, verbose=False) poi = poi.resample_from_to(seg) frame = Snapshot_Frame( img, From 3a9cd6509740387e1d88fd5fb40928d0e92cca69 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 10:57:24 +0000 Subject: [PATCH 20/26] =?UTF-8?q?spinestats:=20version=203=20=E2=80=94=20e?= =?UTF-8?q?ndplate-plane=20Cobb=20+=20apex=20+=20tightened=20multi-Cobb?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Extends the endplate-plane direction machinery used for v2 lordosis/kyphosis to the coronal Cobb angle, and bumps ``CURRENT_VERSION`` to 3 so cached stat files recompute the angle keys (POI buffers stay valid — no new landmarks). Angles (``TPTBox/spine/spinestats/angles.py``): - ``_get_norm`` now routes ``Vertebra_Direction_Right`` through a new ``_endplate_r_direction`` helper (analogue of ``_endplate_ap_direction``), so Cobb lines are true endplate lines averaged across the shared disc instead of the two flanking vertebrae disagreeing on the same boundary. - ``plot_cobb_angle`` always uses that endplate right direction — the ``use_ivd_direction`` branch used ``Vertebra_Disc_Inferior`` from each vertebra separately, giving two different lines at one anatomic disc. - ``compute_angel_between_two_points_`` orders ``id1``/``id2`` by ``Vertebra_Instance.order_dict()`` rather than by numeric label, so T13 (label 28, anatomically between T12 and L1) pairs to the right endplates and ``MoveTo.TOP``/``BOTTOM`` are not swapped. - ``compute_max_cobb_angle_multi`` uses an exclusive split (``+ 1`` on the ``below`` slice) so sibling curves cannot share an endpoint vertebra. - ``VERT_START_COBB`` is C3 (C2 excluded from Cobb search). - ``compute_lordosis_and_kyphosis`` emits ``{curve}_apex`` keys, computed by the new ``_find_curve_apex`` helper (same bisector heuristic as ``compute_max_cobb_angle``). ``plot_compute_lordosis_and_kyphosis`` renders a ``*apex`` marker at that vertebra. - ``plot_cobb_angle`` anchors the label at the apex disc and widens the coronal pad from 50 to 80 mm so labels don't get clipped. Runner (``TPTBox/spine/spinestats/_run_all.py``): - Introduces ``POI_INVALIDATING_VERSION = 2`` so the ``_poi_is_stale_wrt_stat`` threshold is a named constant rather than a magic ``2``, and its docstring points at a symbol that actually exists (previously referenced a nonexistent ``ANGLES_VERSION``). - Documents the v3 bump in the ``CURRENT_VERSION`` history comment. Veridah (``TPTBox/spine/spinestats/veridah_angles.py``): - ``_compute_k_variant`` also returns ``lumbar_lordosis_apex`` so the per-k variants carry the new apex key. Co-Authored-By: Claude Opus 4.7 --- TPTBox/spine/spinestats/_run_all.py | 27 ++- TPTBox/spine/spinestats/angles.py | 198 +++++++++++++++++++--- TPTBox/spine/spinestats/veridah_angles.py | 2 +- 3 files changed, 194 insertions(+), 33 deletions(-) diff --git a/TPTBox/spine/spinestats/_run_all.py b/TPTBox/spine/spinestats/_run_all.py index 0e2598f9..c1d0e1f6 100644 --- a/TPTBox/spine/spinestats/_run_all.py +++ b/TPTBox/spine/spinestats/_run_all.py @@ -39,7 +39,18 @@ # 2 -- endplate-plane based lordosis / kyphosis (average of the two flanking # endplates at each disc), S1 sacrum-endplate landmarks, and correct # cranio-caudal ordering for the T13 annotation label. -CURRENT_VERSION = 2 +# 3 -- endplate-plane based cobb (chain-closed, matches lordosis convention), +# C2 excluded from Cobb search (starts at C3), non-overlapping multi-cobb +# curves, anatomic (not numeric) vertebra ordering in +# ``compute_angel_between_two_points_`` so T13 pairs measure the right +# endplates, and per-curve apex vertebra keys in ``curv`` +# (``{cervical_lordosis,thoracic_kyphosis,lumbar_lordosis}_apex``). +CURRENT_VERSION = 3 +# Highest ``CURRENT_VERSION`` bump that required *new POI landmarks*. Stat files +# older than this were written before those landmarks existed, so the sibling +# POI buffer must be rebuilt (not just the angle keys). Bumps that only change +# how numbers are derived from an unchanged POI set do not update this. +POI_INVALIDATING_VERSION = 2 # Top-level keys we require inside a finished json before we consider a # subject "done" and skip recomputation. cobb/curv are optional and only @@ -162,14 +173,14 @@ def _poi_is_stale_wrt_stat(stat_path: Path, loaded_stat: dict) -> bool: """Return True when the POI buffer sitting next to ``stat_path`` should be rebuilt. Staleness is derived from the stat json's ``_provenance.version``: whenever - that reads < :data:`ANGLES_VERSION`, the sibling POI buffer is treated as - outdated (v1 stat + v1 POI travelled together). ``loaded_stat`` is the - already-loaded stat dict — pass ``{}`` if none exists (then nothing is stale - since there's no v1 marker to invalidate against). + that reads < :data:`POI_INVALIDATING_VERSION`, the sibling POI buffer is + treated as outdated (older stat + older POI travelled together). + ``loaded_stat`` is the already-loaded stat dict — pass ``{}`` if none exists + (then nothing is stale since there's no version marker to invalidate against). """ if not stat_path.exists() or not loaded_stat: return False - return _stat_version(loaded_stat) < CURRENT_VERSION + return _stat_version(loaded_stat) < POI_INVALIDATING_VERSION def _is_cache_valid(json_path: Path, seg_files: list[Path], required_keys: tuple[str, ...]) -> tuple[bool, dict | None]: @@ -391,7 +402,7 @@ def run_all( # invalidating the expensive body-composition / VBQ / muscle_fat / torso # blocks that don't depend on the endplate-based lordosis fix. _cur_ver = _stat_version(out) - if _cur_ver < 2: + if _cur_ver < CURRENT_VERSION: logger.on_warning("version bump", _cur_ver, "->", CURRENT_VERSION, ": recomputing angle keys") for _k in ( "cobb", @@ -1059,7 +1070,7 @@ def _run_one(args: tuple[dict, bool, bool]) -> tuple[str, str | None, dict]: try: if test: subjects = loop_over_repaired_nako(test=True) - total = 15 + total = 10 aggregate = False elif aggregate: subjects = loop_over_repaired_nako(test=False, sort=aggregate) diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 5d703518..2a3934db 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -15,7 +15,7 @@ from TPTBox.spine.snapshot2D.snapshot_modular import Snapshot_Frame, create_snapshot IVD_MORE_ACCURATE = 15 -VERT_START_COBB = Vertebra_Instance.C2 +VERT_START_COBB = Vertebra_Instance.C3 class MoveTo(Enum): @@ -254,8 +254,12 @@ def compute_angel_between_two_points_( return None assert id1 != id2, id1 - # Ensure id1 is less than id2 - if id1 > id2: + # Ensure id1 is anatomically above id2 (not just numerically smaller). + # T13 has label value 28 but sits between T12 (19) and L1 (20) in anatomical + # order; using the label value directly would misplace it and swap the wrong + # MoveTo semantics onto TOP/BOTTOM. + _order = Vertebra_Instance.order_dict() + if _order.get(id1, id1) > _order.get(id2, id2): id1, id2 = id2, id1 # Reorient and rescale the POI data poi.reorient_().rescale_(verbose=False) @@ -388,13 +392,59 @@ def compute_lordosis_and_kyphosis(poi: POI, project_2D=True) -> dict[str, float poi = poi.copy() for k, i in curvature_definition.items(): - angle = compute_angel_between_two_points_( - poi, i.get_start_vert(poi), i.get_stop_vert(poi), "P", i.start_move, i.stop_move, project_2D - ) + start = i.get_start_vert(poi) + stop = i.get_stop_vert(poi) + angle = compute_angel_between_two_points_(poi, start, stop, "P", i.start_move, i.stop_move, project_2D) out[k] = round(angle, 4) if angle is not None else None + out[f"{k}_apex"] = _find_curve_apex(poi, start, stop, i.start_move, i.stop_move, Location.Vertebra_Direction_Posterior) return out +def _find_curve_apex( + poi: POI, + from_vert: Vertebra_Instance | int | None, + to_vert: Vertebra_Instance | int | None, + from_mv: MoveTo, + to_mv: MoveTo, + location: Location, +) -> int | None: + """Return the vertebra between ``from_vert`` and ``to_vert`` whose direction is closest to the endpoint bisector. + + Same apex heuristic as :func:`compute_max_cobb_angle`, just parameterised by + the direction location so it also works for lordosis / kyphosis + (``Vertebra_Direction_Posterior``). Returns ``None`` when either endpoint + direction is unavailable or no intermediate vertebra is present. + """ + if from_vert is None or to_vert is None: + return None + from_v = from_vert if isinstance(from_vert, Vertebra_Instance) else Vertebra_Instance(from_vert) + to_v = to_vert if isinstance(to_vert, Vertebra_Instance) else Vertebra_Instance(to_vert) + a = _get_norm(poi, from_v, from_mv, location, 1) + b = _get_norm(poi, to_v, to_mv, location, 1) + if a is None or b is None: + return None + apex_v = (a + b) / 2 + order = Vertebra_Instance.order() + try: + i_from = order.index(from_v) + i_to = order.index(to_v) + except ValueError: + return None + if i_from > i_to: + i_from, i_to = i_to, i_from + apex: int | None = None + cos_dis = -np.inf + for v in order[i_from : i_to + 1]: + n = _get_norm(poi, v, to_mv, location, 1) + if n is None: + continue + cos_new = cosine_distance(n, apex_v) + if cos_new > cos_dis: + cos_dis = cos_new + apex = v.value + return apex + + def _single_endplate_ap_direction(poi: POI, vert: Vertebra_Instance, side: MoveTo) -> np.ndarray | None: """A/P direction of *one* endplate (superior for TOP, inferior for BOTTOM). @@ -476,6 +526,85 @@ def _endplate_ap_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np. return avg / n +def _single_endplate_r_direction(poi: POI, vert: Vertebra_Instance, side: MoveTo) -> np.ndarray | None: + """R/L direction of *one* endplate (superior for TOP, inferior for BOTTOM). + + Builds a unit vector that lies in the endplate plane along the vertebra's + right/left axis: takes ``corpus - Vertebra_Direction_Right`` (matching the + WK path's ``inv=1`` sign convention, which points *left*) and projects it + onto the plane perpendicular to the endplate normal (``endplate_point - + corpus``). Returns ``None`` if any required POI landmark is missing. + """ + if side == MoveTo.TOP: + endplate_loc = Location.Vertebral_Body_Endplate_Superior + elif side == MoveTo.BOTTOM: + endplate_loc = Location.Vertebral_Body_Endplate_Inferior + else: + return None + if (vert, 50) not in poi or (vert, endplate_loc) not in poi or (vert, Location.Vertebra_Direction_Right) not in poi: + return None + corpus = np.array(poi[vert, 50], dtype=float) + ep = np.array(poi[vert, endplate_loc], dtype=float) + r_pt = np.array(poi[vert, Location.Vertebra_Direction_Right], dtype=float) + n = ep - corpus + if side == MoveTo.BOTTOM: + n = -n # flip inferior endplate so both cases point superior + n_norm = np.linalg.norm(n) + # WK path uses ``corpus - right_pt`` (points anatomical LEFT); mirror that + # convention so both paths agree when they meet in + # ``compute_angel_between_two_points_``. + r_vec = corpus - r_pt + r_norm = np.linalg.norm(r_vec) + if n_norm < 1e-8 or r_norm < 1e-8: + return None + n /= n_norm + r_vec /= r_norm + # Project the vertebra right/left vector onto the endplate plane + r_ep = r_vec - np.dot(r_vec, n) * n + r_ep_norm = np.linalg.norm(r_ep) + if r_ep_norm < 1e-8: + return None + return r_ep / r_ep_norm + + +def _endplate_r_direction(poi: POI, vert: Vertebra_Instance, mv: MoveTo) -> np.ndarray | None: + """R/L direction in the endplate plane at a vertebra-disc *boundary*. + + Analogue of :func:`_endplate_ap_direction` for the coronal (right) direction, + used to obtain a classical endplate-line orientation for Cobb angles. Averages + the two flanking endplate R directions at a disc: + + - :attr:`MoveTo.BOTTOM` at vertebra ``V`` averages ``V``'s inferior endplate + with the superior endplate of the next vertebra. + - :attr:`MoveTo.TOP` at vertebra ``V`` averages ``V``'s superior endplate + with the inferior endplate of the previous vertebra. + + Falls back to the single-endplate direction when the neighbour is missing. + Returns ``None`` when even the local endplate is unavailable. + """ + if isinstance(vert, int): + vert = Vertebra_Instance(vert) + if mv == MoveTo.BOTTOM: + neighbour = vert.get_next_poi(poi) + neighbour_side = MoveTo.TOP + elif mv == MoveTo.TOP: + neighbour = vert.get_previous_poi(poi) + neighbour_side = MoveTo.BOTTOM + else: + return _single_endplate_r_direction(poi, vert, mv) + own = _single_endplate_r_direction(poi, vert, mv) + other = _single_endplate_r_direction(poi, neighbour, neighbour_side) if neighbour is not None else None + if own is None: + return other + if other is None: + return own + avg = own + other + n = np.linalg.norm(avg) + if n < 1e-8: + return own + return avg / n + + def _get_norm(poi: POI, id1: int | Vertebra_Instance, mv: MoveTo, location: Location, inv: int = 1) -> np.ndarray | None: # noqa: ARG001 """Return the normalised direction vector from a location POI to the vertebra centroid. @@ -484,7 +613,10 @@ def _get_norm(poi: POI, id1: int | Vertebra_Instance, mv: MoveTo, location: Loca landmark (``Vertebral_Body_Endplate_Superior`` / ``_Inferior``) is preferred over the averaged vertebral-body posterior direction — this yields the classical endplate-line orientation used in Cobb-style lordosis/kyphosis measurements. - The endplate direction is used only when both the relevant endplate point and + The same automatic switch applies to ``Location.Vertebra_Direction_Right``: it + is projected into the endplate plane so Cobb (scoliosis) angles are measured + between endplate lines, analogous to the sagittal case. The endplate direction + is used only when both the relevant endplate point and ``Vertebra_Direction_Right`` are present in ``poi`` for that vertebra; otherwise the code falls back to the WK-based averaged direction below. """ @@ -494,6 +626,10 @@ def _get_norm(poi: POI, id1: int | Vertebra_Instance, mv: MoveTo, location: Loca ep_norm = _endplate_ap_direction(poi, id1, mv) if ep_norm is not None: return ep_norm * inv + if location == Location.Vertebra_Direction_Right and mv in (MoveTo.TOP, MoveTo.BOTTOM): + ep_norm = _endplate_r_direction(poi, id1, mv) + if ep_norm is not None: + return ep_norm * inv subreg = 50 if location in [Location.Vertebra_Disc_Inferior, Location.Vertebra_Disc_Superior]: subreg = 100 @@ -710,8 +846,13 @@ def compute_max_cobb_angle_multi( assert from_vert is not None assert to_vert is not None out_list.append((max_angle, from_vert, to_vert, apex)) + # Exclusive split: neither endpoint of the just-found curve may participate + # in a subsequent curve. Prevents overlaps like (T1-T5) + (T5-T7); a + # sibling curve below the current one starts strictly caudal to to_vert, + # a sibling above ends strictly cranial to from_vert. Drop the ``+ 1`` + # on the ``below`` slice to restore the textbook (endpoint-shared) split. above = vertebrae_list[: vertebrae_list.index(Vertebra_Instance(from_vert))] - below = vertebrae_list[vertebrae_list.index(Vertebra_Instance(to_vert)) :] + below = vertebrae_list[vertebrae_list.index(Vertebra_Instance(to_vert)) + 1 :] compute_max_cobb_angle_multi( poi, above, @@ -846,8 +987,15 @@ def plot_compute_lordosis_and_kyphosis( out.append((id1.value, s, (-a[0] * line_len * 3, -a[1] * line_len * 3))) out2 = compute_lordosis_and_kyphosis(poi, project_2D=project_2D) for name, v in out2.items(): - if v is None: + if v is None or name not in curvature_definition: + # Skip auxiliary keys like ``*_apex`` that live alongside the angles + # in the same dict but have no curve definition of their own. continue + # Apex annotation: mark the apex vertebra body with a star + label so + # the reader can see which vertebra the ``*_apex`` json key refers to. + apex_v = out2.get(f"{name}_apex") + if apex_v is not None and (apex_v, 50) in poi: + text_out.append((apex_v, ("*apex", -60))) id1 = curvature_definition[name].get_start_vert(poi) id2 = curvature_definition[name].get_stop_vert(poi) @@ -935,32 +1083,34 @@ def plot_cobb_angle( for id1, mv in zip([from_vert, to_vert], [vert_id1_mv, vert_id2_mv]): c = mv.get_location(id1, poi) - if use_ivd_direction and id1 > IVD_MORE_ACCURATE: - norm1_post = _get_norm(poi, id1, mv, Location.Vertebra_Direction_Posterior) - a = _get_norm(poi, id1, mv, Location.Vertebra_Disc_Inferior) - a = np.cross(a, norm1_post) - else: - a = _get_norm(poi, id1, mv, Location.Vertebra_Direction_Right) - - # print(a, id1, mv, c) + # Always go through Vertebra_Direction_Right so _get_norm routes + # to the endplate-plane right direction (chain-closed across the + # shared disc: `_endplate_r_direction` averages both flanking + # endplates). The old ``use_ivd_direction`` branch pulled + # ``Vertebra_Disc_Inferior`` from each vertebra separately — + # T9-BOTTOM used the T9/T10 disc but T10-TOP used the T10/T11 + # disc, so the same anatomic boundary got two different lines. + a = _get_norm(poi, id1, mv, Location.Vertebra_Direction_Right) assert a is not None out.append((apex, c, (-a[2] * line_len, a[1] * line_len))) out.append((apex, c, (a[2] * line_len, -a[1] * line_len))) - # a = _get_norm(poi, id1, mv, Location.Vertebra_Disc_Inferior) - # out.append((apex, c, (a[2] * line_len, -a[1] * line_len))) if apex is not None: - cord = poi[apex, 50] - s = f"copp angle: {max_angle:.1f}° {Vertebra_Instance(from_vert)} - {Vertebra_Instance(to_vert)}" - text_out.append((apex, (s, 25, cord[1]))) + # Align the label with the disc below the apex vertebra (IVD height) + # rather than the vertebra body centre, so the text sits at the same + # cranio-caudal level as the drawn Cobb line at the apex. + cord = poi[apex, Location.Vertebra_Disc.value] if (apex, Location.Vertebra_Disc.value) in poi else poi[apex, 50] + s = f"copp angle\n{max_angle:.1f}° {Vertebra_Instance(from_vert)} - {Vertebra_Instance(to_vert)}" + text_out.append((apex, (s, 35, cord[1]))) poi.info["line_segments_cor"] = out + poi.info.get("line_segments_cor", []) poi.info["text_cor"] = text_out + poi.info.get("text_cor", []) axis = poi.get_axis("R") width = poi.shape[axis] / poi.zoom[axis] / 2 - if width < 50: + min_half_width_mm = 80 + if width < min_half_width_mm: padd = [(0, 0) for _ in range(3)] - padd[axis] = (int(50 - width), int(50 - width)) + padd[axis] = (int(min_half_width_mm - width), int(min_half_width_mm - width)) img = to_nii(img).apply_pad(padd, verbose=False) seg = to_nii(seg, True).apply_pad(padd, verbose=False) poi = poi.resample_from_to(seg) diff --git a/TPTBox/spine/spinestats/veridah_angles.py b/TPTBox/spine/spinestats/veridah_angles.py index a607319c..2fa366d7 100644 --- a/TPTBox/spine/spinestats/veridah_angles.py +++ b/TPTBox/spine/spinestats/veridah_angles.py @@ -97,7 +97,7 @@ def _compute_k_variant(poi: POI, k: int, project_2D: bool = False) -> dict[str, if not relabeled.centroids: return {"lumbar_lordosis": None} full = compute_lordosis_and_kyphosis(relabeled, project_2D=project_2D) - return {"lumbar_lordosis": full.get("lumbar_lordosis")} + return {"lumbar_lordosis": full.get("lumbar_lordosis"), "lumbar_lordosis_apex": full.get("lumbar_lordosis_apex")} def _veridah_region_map(veridah_json_path: Path) -> dict[int, int] | None: From c8ccbce40b762ceb918db425061d83cf93f8524f Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 11:10:19 +0000 Subject: [PATCH 21/26] extended dixon mapping MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Cherry-picked content from tanja@50c8b4d without the whitespace-only reformatting noise: - Add h2d/tse2d1-4/tser2d regex to series-description → format mapping; drop the broad ".*t2w?_tse.*" match. - Fall back to SequenceName when SeriesDescription is empty. Co-Authored-By: Tanja Lerchl Co-Authored-By: Claude Opus 4.7 --- TPTBox/core/dicom/dicom_header_to_keys.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/TPTBox/core/dicom/dicom_header_to_keys.py b/TPTBox/core/dicom/dicom_header_to_keys.py index 54db4da4..8f483e6b 100644 --- a/TPTBox/core/dicom/dicom_header_to_keys.py +++ b/TPTBox/core/dicom/dicom_header_to_keys.py @@ -45,7 +45,11 @@ } dixon_mapping = {**dixon_mapping, **{v: v for v in dixon_mapping.values()}} map_series_description_to_file_format_default = { - ".*t2w?_tse.*": "T2w", + ".*h2d.*": "T2haste", + ".*tse2d1-4.*": "T1w", + ".*tser2d.*": "T2w", + # ".*tse2d1_4.*": "T1w", + # ".*_tse.*": "T2w", "t2w?_fse.*": "T2w", ".*t1w?_tse.*": "T1w", ".*t1w?_vibe_tra.*": "vibe", @@ -297,6 +301,11 @@ def _get(key, default=None): keys["ce"] = "ContrastAgent" # GET MRI FORMAT series_description = _get("SeriesDescription", "mr").lower() + if series_description == "mr": + series_description = _get("SequenceName", "mr").lower() + print( + f"SeriesDescription: '{series_description}', ImageType: {image_type}, ProtocolName: '{_get('ProtocolName', '')}'" + ) modality = _get("Modality", "mr").lower() mri_format = None From 2a51ed673778fdfd62af66b84cf69cc7295f2a25 Mon Sep 17 00:00:00 2001 From: robert-graf <31210726+robert-graf@users.noreply.github.com> Date: Tue, 22 Sep 2026 11:13:23 +0000 Subject: [PATCH 22/26] style fixes by ruff --- TPTBox/core/dicom/dicom_header_to_keys.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/TPTBox/core/dicom/dicom_header_to_keys.py b/TPTBox/core/dicom/dicom_header_to_keys.py index 8f483e6b..0aa37c36 100644 --- a/TPTBox/core/dicom/dicom_header_to_keys.py +++ b/TPTBox/core/dicom/dicom_header_to_keys.py @@ -303,9 +303,7 @@ def _get(key, default=None): series_description = _get("SeriesDescription", "mr").lower() if series_description == "mr": series_description = _get("SequenceName", "mr").lower() - print( - f"SeriesDescription: '{series_description}', ImageType: {image_type}, ProtocolName: '{_get('ProtocolName', '')}'" - ) + print(f"SeriesDescription: '{series_description}', ImageType: {image_type}, ProtocolName: '{_get('ProtocolName', '')}'") modality = _get("Modality", "mr").lower() mri_format = None From 01379ece805f9e23b3a5c12011195b94427a8935 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 11:18:59 +0000 Subject: [PATCH 23/26] style: ruff --fix + ruff format across the branch MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Manual fixes in ``TPTBox/spine/spinestats/_load_nako_wh.py``: - D205 (verify_missing_images, verify_hardlink): reword docstrings so the summary line stands alone from the description body. - PLW2901 (``v = preferred`` at line 787): add ``# noqa: PLW2901`` to match the existing noqa on the sibling ``k = mapping.get(k, k)`` rebinding — the loop deliberately overrides both bindings. - PERF102 (line 1385): iterate the ``t2w_chunk`` mapping via ``.values()`` since only the value is used. - TRY300 (``_grid_worker``): move the success ``return`` out of the ``try`` block; only ``_add_grid_info_to_json`` may raise. Everything else is ``ruff check --fix`` and ``ruff format`` output. Co-Authored-By: Claude Opus 4.7 --- TPTBox/core/dicom/dicom_extract.py | 6 +-- TPTBox/core/internal/train_nnUnet/train.py | 8 +-- TPTBox/core/poi_fun/vector_fields.py | 2 +- TPTBox/spine/spinestats/_load_nako_wh.py | 21 ++++---- TPTBox/spine/spinestats/_qc_report.py | 37 +++++++------- .../measure_ivd_and_vertebra_geometry.py | 12 ++--- docs/api/poi_fun.md | 8 +-- unit_tests/test_poi_vector_fields.py | 51 +++++++++++++++---- 8 files changed, 85 insertions(+), 60 deletions(-) diff --git a/TPTBox/core/dicom/dicom_extract.py b/TPTBox/core/dicom/dicom_extract.py index 12eeb171..943f5ffe 100644 --- a/TPTBox/core/dicom/dicom_extract.py +++ b/TPTBox/core/dicom/dicom_extract.py @@ -1031,11 +1031,7 @@ def extract_dicom_folder( # → no DICOM headers read for this source. Folders fingerprint their # rglob'd file list; zips fingerprint (size, mtime_ns) of the archive # itself (see `_source_fingerprint`). - if ( - skip_already_extracted - and not force_rescan - and _is_already_extracted(Path(dicom_path), Path(dataset_path_out)) - ): + if skip_already_extracted and not force_rescan and _is_already_extracted(Path(dicom_path), Path(dataset_path_out)): logger.print(f"Skip {dicom_path} (already extracted; fingerprint matches)", verbose=verbose) continue # Track the original source path so the marker below is keyed to the diff --git a/TPTBox/core/internal/train_nnUnet/train.py b/TPTBox/core/internal/train_nnUnet/train.py index 25ee589a..a9677dc4 100644 --- a/TPTBox/core/internal/train_nnUnet/train.py +++ b/TPTBox/core/internal/train_nnUnet/train.py @@ -190,9 +190,7 @@ def _ensure_smauglab_trainer_installed(trainer_class_name: str) -> None: f"`sudo smauglab_add_nnunettrainer --trainer {module_name} --overwrite` once." ) from e except OSError as e: - raise RuntimeError( - f"Failed to install SmaugLab trainer module {module_name!r} into nnunetv2 at {dst}: {e}" - ) from e + raise RuntimeError(f"Failed to install SmaugLab trainer module {module_name!r} into nnunetv2 at {dst}: {e}") from e def _apply_smauglab_params_env(trainer_class_name: str, dataset_folder: Path) -> None: @@ -278,9 +276,7 @@ def _run_training( assert not disable_checkpointing, "--val_best is not compatible with --disable_checkpointing" try: - nnunet_trainer = get_trainer_from_args( - dataset_name_or_id, configuration, fold, trainer_class_name, plans_identifier, device=device - ) + nnunet_trainer = get_trainer_from_args(dataset_name_or_id, configuration, fold, trainer_class_name, plans_identifier, device=device) except RuntimeError as e: hint = "" if trainer_class_name in _SMAUGLAB_TRAINERS: diff --git a/TPTBox/core/poi_fun/vector_fields.py b/TPTBox/core/poi_fun/vector_fields.py index 078cc693..e3271d4a 100644 --- a/TPTBox/core/poi_fun/vector_fields.py +++ b/TPTBox/core/poi_fun/vector_fields.py @@ -107,7 +107,7 @@ def _remap_vector_field_keys_inplace(info: dict, region_map: dict | None, subreg Keys may be integer labels or ``Vertebra_Instance``-name strings; both are matched against the int-keyed maps. No-op if there is nothing to do. """ - from TPTBox.core.poi_fun.poi_abstract import LABEL_NAME, _GROUP_NAME_KEY, label_name_dict + from TPTBox.core.poi_fun.poi_abstract import _GROUP_NAME_KEY, LABEL_NAME, label_name_dict if not region_map and not subregion_map: return diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index fa18c64d..6746b9cf 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -199,8 +199,9 @@ def resolve_pick( def verify_missing_images(cache: DecisionCache, sub: str, subj_dict: dict) -> None: - """For each expected base image absent from ``subj_dict``, ask the user whether it's - really missing. If confirmed missing, drop the base and its dependent seg keys from + """For each expected base image absent from ``subj_dict``, ask the user whether it's really missing. + + If confirmed missing, drop the base and its dependent seg keys from ``subj_dict`` (set to None). Decisions are cached per (subject, image). """ for base, deps in EXPECTED_IMAGES.items(): @@ -784,7 +785,7 @@ def loop_over_repaired_nako( if k in _MEVIBE_EXTRA_KEYS and len(v) > 1: preferred = [bf for bf in v if "/derivatives_mevibe/" in _fmt_file(bf)] if preferred: - v = preferred + v = preferred # noqa: PLW2901 k = mapping.get(k, k) # noqa: PLW2901 if len(v) > 1: picked = resolve_pick(cache, sub, f"mevibe:{k}", f"Multiple mevibe files for {k}; pick one.", v) @@ -1344,10 +1345,12 @@ def _apply_corrections_to_subj_dict(sub: str, subj_dict: dict, index: dict) -> d def verify_hardlink(sub: str, corrected_index_path: Path = _DEFAULT_CORRECTED_INDEX) -> None: - """Run the loop for a single subject, apply corrections, then trace what - :func:`hard_link` *would* do and check every source exists and every target - would land on the same filesystem as its source (so :func:`os.link` won't - hit ``EXDEV``). Prints one line per file — no writes are performed. + """Run the loop for a single subject and verify what :func:`hard_link` would do. + + Applies corrections, traces what :func:`hard_link` *would* do, and checks + every source exists and every target would land on the same filesystem as + its source (so :func:`os.link` won't hit ``EXDEV``). Prints one line per + file — no writes are performed. """ index = load_corrected_index(corrected_index_path) sub = str(sub) @@ -1382,7 +1385,7 @@ def _check(bf, parent: str, info: dict | None = None) -> None: status = "ok" if src_exists and same_fs else ("cross-fs" if src_exists else "src-missing") print(f" [{status:>10s}] {src} -> {target}") - for _key, t2w in (d.get("t2w_chunk") or {}).items(): + for t2w in (d.get("t2w_chunk") or {}).values(): if t2w: _check(t2w[0], parent="rawdata", info={"ses": "baseline"}) seg_keys = ( @@ -1433,9 +1436,9 @@ def _grid_worker(nii_path: str) -> tuple[str, str]: sidecar = Path(str(p).split(".")[0] + ".json") try: _add_grid_info_to_json(p, sidecar, add=True) - return nii_path, "ok" except Exception as e: # noqa: BLE001 return nii_path, f"error: {type(e).__name__}: {e}" + return nii_path, "ok" def _iter_grid_targets(subj_dict: dict): diff --git a/TPTBox/spine/spinestats/_qc_report.py b/TPTBox/spine/spinestats/_qc_report.py index fd5eb02a..54d3b122 100644 --- a/TPTBox/spine/spinestats/_qc_report.py +++ b/TPTBox/spine/spinestats/_qc_report.py @@ -79,9 +79,18 @@ def _summary_row(df: pd.DataFrame, col: str, lo: float, hi: float) -> dict: n_out = int(mask_out.sum()) if n_valid == 0: return { - "column": col, "n_total": n_total, "n_valid": 0, "n_missing": n_total, - "median": None, "iqr_low": None, "iqr_high": None, "min": None, "max": None, - "outlier_range": f"[{lo}, {hi}]", "n_outliers": 0, "pct_outliers": None, + "column": col, + "n_total": n_total, + "n_valid": 0, + "n_missing": n_total, + "median": None, + "iqr_low": None, + "iqr_high": None, + "min": None, + "max": None, + "outlier_range": f"[{lo}, {hi}]", + "n_outliers": 0, + "pct_outliers": None, } return { "column": col, @@ -99,9 +108,7 @@ def _summary_row(df: pd.DataFrame, col: str, lo: float, hi: float) -> dict: } -def _outlier_frame( - df: pd.DataFrame, col: str, lo: float, hi: float, extra_cols: list[str] -) -> pd.DataFrame: +def _outlier_frame(df: pd.DataFrame, col: str, lo: float, hi: float, extra_cols: list[str]) -> pd.DataFrame: v = pd.to_numeric(df[col], errors="coerce") mask = v.notna() & ((v < lo) | (v > hi)) if not mask.any(): @@ -133,8 +140,10 @@ def build_qc_report(folder: Path) -> Path: {"item": "per_subject_cols", "value": len(sub.columns)}, {"item": "n_vertebra_rows", "value": len(vert)}, {"item": "n_ivd_rows", "value": len(ivd)}, - {"item": "pelvic_error_rate_%", - "value": round(100.0 * sub.get("pelvic_parameters.error", pd.Series([np.nan] * len(sub))).notna().sum() / len(sub), 3)}, + { + "item": "pelvic_error_rate_%", + "value": round(100.0 * sub.get("pelvic_parameters.error", pd.Series([np.nan] * len(sub))).notna().sum() / len(sub), 3), + }, ] ) @@ -163,15 +172,9 @@ def build_qc_report(folder: Path) -> Path: # ------------------------------------------------------------------ # Per-column summary # ------------------------------------------------------------------ - sub_summary = pd.DataFrame( - [_summary_row(sub, c, lo, hi) for c, (lo, hi) in OUTLIER_SUBJECT.items() if c in sub.columns] - ) - vert_summary = pd.DataFrame( - [_summary_row(vert, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in vert.columns] - ) - ivd_summary = pd.DataFrame( - [_summary_row(ivd, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in ivd.columns] - ) + sub_summary = pd.DataFrame([_summary_row(sub, c, lo, hi) for c, (lo, hi) in OUTLIER_SUBJECT.items() if c in sub.columns]) + vert_summary = pd.DataFrame([_summary_row(vert, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in vert.columns]) + ivd_summary = pd.DataFrame([_summary_row(ivd, c, lo, hi) for c, (lo, hi) in OUTLIER_LABEL.items() if c in ivd.columns]) # ------------------------------------------------------------------ # Outlier rows diff --git a/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py index 163a72ba..1d08bae6 100644 --- a/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py +++ b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py @@ -197,9 +197,7 @@ def measure_ivd_and_vertebra_geometry( raw = _compute_directional_heights_widths(vert, poi, label, step_size_mm=step_size_mm, raw=raw) # 3. normalized T2 signal if t2w_arr is not None: - raw = _compute_t2_signal_ratio( - t2w_arr, vert, label, spinal_canal_signal, spinal_canal_signal_old, raw=raw, erode=erode - ) + raw = _compute_t2_signal_ratio(t2w_arr, vert, label, spinal_canal_signal, spinal_canal_signal_old, raw=raw, erode=erode) results[label] = _result_from_info(info) except Exception as e: results[label] = _nan_result(error=str(e)) @@ -459,9 +457,7 @@ def _batched_ray_segments(mesh: trimesh.Trimesh, ray_direction: np.ndarray, orig """ n = origins.shape[0] directions = np.broadcast_to(ray_direction, (n, 3)) - locations, index_ray, _ = mesh.ray.intersects_location( - ray_origins=origins, ray_directions=directions, multiple_hits=True - ) + locations, index_ray, _ = mesh.ray.intersects_location(ray_origins=origins, ray_directions=directions, multiple_hits=True) lengths = np.zeros(n) first_pts = np.zeros((n, 3)) last_pts = np.zeros((n, 3)) @@ -484,7 +480,9 @@ def _batched_ray_segments(mesh: trimesh.Trimesh, ray_direction: np.ndarray, orig return lengths, first_pts, last_pts -def _grid_origins(base: np.ndarray, v1: np.ndarray, v2: np.ndarray, xs: np.ndarray, ys: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]: +def _grid_origins( + base: np.ndarray, v1: np.ndarray, v2: np.ndarray, xs: np.ndarray, ys: np.ndarray +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Return (origins, flat_x, flat_y) for a 2D grid of ray origins in the (v1, v2) plane.""" grid_x, grid_y = np.meshgrid(xs, ys, indexing="ij") flat_x = grid_x.ravel() diff --git a/docs/api/poi_fun.md b/docs/api/poi_fun.md index b4cd34ae..f81f5ae5 100644 --- a/docs/api/poi_fun.md +++ b/docs/api/poi_fun.md @@ -99,11 +99,11 @@ Use the accessors on `Abstract_POI` (available on both `POI` and `POI_Global`) instead of writing the dict directly: ```python -poi.set_label_name(region=2, subregion=10, name="FLCPC") # per-point label -poi.set_level_one_name(region=2, name="Femur") # region group name +poi.set_label_name(region=2, subregion=10, name="FLCPC") # per-point label +poi.set_level_one_name(region=2, name="Femur") # region group name -poi.label_name(2, 10) # -> "FLCPC" -poi.level_one_name(2) # -> "Femur" +poi.label_name(2, 10) # -> "FLCPC" +poi.level_one_name(2) # -> "Femur" ``` `region` / `subregion` accept `int`, numeric string, or `Enum` members. A diff --git a/unit_tests/test_poi_vector_fields.py b/unit_tests/test_poi_vector_fields.py index cac5fdb6..98922511 100644 --- a/unit_tests/test_poi_vector_fields.py +++ b/unit_tests/test_poi_vector_fields.py @@ -269,16 +269,44 @@ class Test_LabelName_LegacyMigration(unittest.TestCase): """ _ATLAS_FLAT: ClassVar[dict[str, str]] = { - "(1, 1)": "TGT", "(1, 2)": "FHC", "(1, 3)": "FNC", "(1, 4)": "FAAP", - "(2, 1)": "FLCD", "(2, 2)": "FMCD", "(2, 3)": "FLCP", "(2, 4)": "FMCP", - "(2, 5)": "FNP", "(2, 6)": "FADP", "(2, 7)": "TGPP", "(2, 8)": "TGCP", - "(2, 9)": "FMCPC", "(2, 10)": "FLCPC", "(2, 11)": "TRMP", "(2, 12)": "TRLP", - "(3, 1)": "TLCL", "(3, 2)": "TMCM", "(3, 3)": "TKC", "(3, 4)": "TLCA", - "(3, 5)": "TLCP", "(3, 6)": "TMCA", "(3, 7)": "TMCP", "(3, 8)": "TTP", - "(3, 9)": "TAAP", "(3, 10)": "TMIT", "(3, 11)": "TLIT", - "(4, 1)": "FLM", "(4, 2)": "TMM", "(4, 3)": "TAC", "(4, 4)": "TADP", - "(5, 1)": "PPP", "(5, 2)": "PDP", "(5, 3)": "PMP", "(5, 4)": "PLP", - "(5, 5)": "PRPP", "(5, 6)": "PRDP", "(5, 7)": "PRHP", + "(1, 1)": "TGT", + "(1, 2)": "FHC", + "(1, 3)": "FNC", + "(1, 4)": "FAAP", + "(2, 1)": "FLCD", + "(2, 2)": "FMCD", + "(2, 3)": "FLCP", + "(2, 4)": "FMCP", + "(2, 5)": "FNP", + "(2, 6)": "FADP", + "(2, 7)": "TGPP", + "(2, 8)": "TGCP", + "(2, 9)": "FMCPC", + "(2, 10)": "FLCPC", + "(2, 11)": "TRMP", + "(2, 12)": "TRLP", + "(3, 1)": "TLCL", + "(3, 2)": "TMCM", + "(3, 3)": "TKC", + "(3, 4)": "TLCA", + "(3, 5)": "TLCP", + "(3, 6)": "TMCA", + "(3, 7)": "TMCP", + "(3, 8)": "TTP", + "(3, 9)": "TAAP", + "(3, 10)": "TMIT", + "(3, 11)": "TLIT", + "(4, 1)": "FLM", + "(4, 2)": "TMM", + "(4, 3)": "TAC", + "(4, 4)": "TADP", + "(5, 1)": "PPP", + "(5, 2)": "PDP", + "(5, 3)": "PMP", + "(5, 4)": "PLP", + "(5, 5)": "PRPP", + "(5, 6)": "PRDP", + "(5, 7)": "PRHP", } def test_normalize_label_name_migrates_flat_atlas(self): @@ -376,7 +404,8 @@ def test_migration_then_map_labels(self): class Test_LabelName_Accessors(unittest.TestCase): """`set_label_name` / `set_level_one_name` write into ``info['label_name']`` and get read back by ``label_name`` / ``level_one_name`` and by the - Slicer/mkr exporter.""" + Slicer/mkr exporter. + """ def _poi_with_enums(self) -> POI: from TPTBox.core.vert_constants import Location, Vertebra_Instance From 1a723a803ebc513c187ef43d16c1f700d1196847 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 11:33:17 +0000 Subject: [PATCH 24/26] fix example test case compile --- TPTBox/spine/spinestats/_load_nako_wh.py | 14 ++++++++++++-- docs/api/poi_fun.md | 4 ++-- 2 files changed, 14 insertions(+), 4 deletions(-) diff --git a/TPTBox/spine/spinestats/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py index 6746b9cf..dd739636 100644 --- a/TPTBox/spine/spinestats/_load_nako_wh.py +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -761,8 +761,18 @@ def loop_over_repaired_nako( q.filter_format("mevibe") # q.filter("sequ", "me1") mevibe_fams = list(q.loop_dict(key_addendum=["mod", "part", "desc"])) - # Drop derivative-only families that don't carry the raw echo images. - mevibe_fams = [f for f in mevibe_fams if "mevibe_part-eco0-opp1" in f] + # Drop derivative-only families and incomplete acquisitions: get_corrected_mevibe + # unconditionally indexes all six echoes, so a family missing any of them would crash. + _echo_keys = [f"mevibe_part-{k}" for k in ("eco0-opp1", "eco1-pip1", "eco2-opp2", "eco3-in1", "eco4-pop1", "eco5-arb1")] + _kept = [] + for f in mevibe_fams: + missing = [k for k in _echo_keys if k not in f] + if missing: + if "mevibe_part-eco0-opp1" in f: + log.on_warning(f"sub-{sub}: incomplete mevibe family {f.family_id!r} (missing {missing}); skipping") + continue + _kept.append(f) + mevibe_fams = _kept if len(mevibe_fams) > 1: labels = [str(f.get("mevibe_part-eco0-opp1", f)) for f in mevibe_fams] cached_pick = _cached_pick(cache.get(sub, "mevibe_fam")) diff --git a/docs/api/poi_fun.md b/docs/api/poi_fun.md index f81f5ae5..675c4fb2 100644 --- a/docs/api/poi_fun.md +++ b/docs/api/poi_fun.md @@ -82,13 +82,13 @@ from TPTBox.core.poi_fun.vector_fields import ( vec_fields = poi.info.setdefault(POI_INFO_VECTOR_FIELDS_KEY, []) if "my_vector_field" not in vec_fields: vec_fields.append("my_vector_field") -poi.info["my_vector_field"] = {"L1": (0.09, -0.99, -0.02), ...} +poi.info["my_vector_field"] = {"L1": (0.09, -0.99, -0.02)} # ... # Scalar per-vertebra metadata (key-remapped by map_labels only): lbl_fields = poi.info.setdefault(POI_INFO_LABEL_KEYED_FIELDS_KEY, []) if "my_scalar_field" not in lbl_fields: lbl_fields.append("my_scalar_field") -poi.info["my_scalar_field"] = {"L1": 3.14, ...} +poi.info["my_scalar_field"] = {"L1": 3.14} # ... ``` **Assigning names (`info["label_name"]`):** From cbf9afa60e597080738c0b78820ef94fa2098677 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 12:08:21 +0000 Subject: [PATCH 25/26] skipt T13 for apex --- TPTBox/spine/spinestats/angles.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 2a3934db..2ac35536 100644 --- a/TPTBox/spine/spinestats/angles.py +++ b/TPTBox/spine/spinestats/angles.py @@ -751,6 +751,8 @@ def compute_max_cobb_angle( assert b is not None apex_v = (a + b) / 2 for i in vertebrae_list[vertebrae_list.index(Vertebra_Instance(from_vert)) : vertebrae_list.index(Vertebra_Instance(to_vert)) + 1]: + if i.value not in poi.keys_region(): + continue try: a = _get_norm(poi, i, vert_id2_mv, Location.Vertebra_Direction_Right, 1) if a is None: From d58bfd027040de97c4648e18b9a84068bea5df86 Mon Sep 17 00:00:00 2001 From: Robert Graf Date: Tue, 22 Sep 2026 12:27:42 +0000 Subject: [PATCH 26/26] last_vert must be in poi --- TPTBox/spine/spinestats/poi_fun/endplates.py | 37 ++++++++++++-------- 1 file changed, 23 insertions(+), 14 deletions(-) diff --git a/TPTBox/spine/spinestats/poi_fun/endplates.py b/TPTBox/spine/spinestats/poi_fun/endplates.py index b14ec2bd..b2dab43d 100644 --- a/TPTBox/spine/spinestats/poi_fun/endplates.py +++ b/TPTBox/spine/spinestats/poi_fun/endplates.py @@ -454,9 +454,16 @@ def calc_endplate_points_( endplate_nii = c.extract_label(superior_label) cms_local_override = None - last_vert = max(vert_ids) - - if (last_vert, Location.Vertebral_Body_Endplate_Inferior.value) in poi: + # A vertebra can appear in the segmentation (vert_ids) but have no POI + # entry (e.g. too few voxels for a centroid). Pick the highest vertebra + # that actually has a centroid in poi so the fallback to Vertebra_Corpus + # below cannot KeyError. + vert_ids_in_poi = [v for v in vert_ids if (v, Location.Vertebra_Corpus.value) in poi] + last_vert = max(vert_ids_in_poi) if vert_ids_in_poi else None + + if last_vert is None: + pass # no anchor available; skip the sacrum endplate override + elif (last_vert, Location.Vertebral_Body_Endplate_Inferior.value) in poi: cms_local_override = poi[last_vert, Location.Vertebral_Body_Endplate_Inferior] elif (last_vert, Location.Vertebra_Disc.value) in poi: cms_local_override = poi[last_vert, Location.Vertebra_Disc.value] @@ -464,17 +471,19 @@ def calc_endplate_points_( cms_local_override = vert.extract_label(100 + last_vert).center_of_masses()[1] else: cms_local_override = poi[last_vert, Location.Vertebra_Corpus] - _endplate( - poi, - endplate_nii, - Location.Vertebral_Body_Endplate_Superior, - Vertebra_Instance.S1.value, - log, - normals_by_vert, - cms_local_override=cms_local_override, - flip_direction=True, - compute_curvature=compute_curvature, - ) + + if last_vert is not None: + _endplate( + poi, + endplate_nii, + Location.Vertebral_Body_Endplate_Superior, + Vertebra_Instance.S1.value, + log, + normals_by_vert, + cms_local_override=cms_local_override, + flip_direction=True, + compute_curvature=compute_curvature, + ) # Angle between superior and inferior endplate normals, per vertebra. for vert_id, normals in normals_by_vert.items(): n_sup = normals.get(Location.Vertebral_Body_Endplate_Superior)