diff --git a/CLAUDE.md b/CLAUDE.md index 87c77ae..aa87d99 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 0000000..0f578ab --- /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`. diff --git a/TPTBox/core/bids_constants.py b/TPTBox/core/bids_constants.py index d534de3..f5a6a08 100755 --- a/TPTBox/core/bids_constants.py +++ b/TPTBox/core/bids_constants.py @@ -129,6 +129,9 @@ "XPC", "phot", "TOF", # Time-of-flight + "MIP", # Maximum Intensity Projection + "DIR", # Double Inversion Recovery + "xa-helper", # non-imaging XA payload (SECONDARY / REFIMAGE / EXAM PROTOCOL) "NerveVIEW", # https://www.philips.de/healthcare/product/HCNMRB971/3D-NerveVIEW-Klinische-MR-Anwendung "3DDrive", # https://www.philips.de/healthcare/product/HCNMRB178/3D-DRIVE-MR-Software "DCE", # dynamic contrast-enhanced () " @@ -165,6 +168,39 @@ "labels", "report", "pet", + # Non-MR imaging modalities handled by `extract_keys_from_json`'s modality + # fallback (see `dicom_header_to_keys.py`). Listed here so `BIDS_FILE` + # accepts them without `non_strict_mode`. + "xray", # 2D X-ray family: CR / DX / RG / MG / PX / IO + "us", # ultrasound + "nm", # nuclear medicine (planar / SPECT) + "sc", # secondary capture (screenshots, derived stills) + "photo", # ophthalmic / external / visible-light photography (OP / XC) + "endoscopy", # ES + "rtimage", # RT image + "rtstruct", # RT structure set + "rtdose", # RT dose grid + "rtplan", # RT plan + "ot", # explicit DICOM "Other" modality + # Non-imaging metadata DICOMs — routed to `.txt` sidecars by the modality + # fallback in `extract_keys_from_json` because they carry no pixel volume. + "pr", # presentation state (GSPS / CSPS) + "ko", # key object selection + "reg", # registration + "fid", # fiducials + "rwv", # real world value map + "plan", # plan + "stain", # automated slide stainer + "resp", # respiratory waveform + "hd", # hemodynamic waveform + "ecg", # electrocardiography + "eps", # cardiac electrophysiology + "ar", # autorefraction + "ker", # keratometry + "len", # lensometry + "va", # visual acuity + "opv", # ophthalmic visual field + "opm", # ophthalmic mapping ] # https://bids-specification.readthedocs.io/en/stable/appendices/entity-table.html formats_relaxed = [*formats, "t2", "t1", "t2c", "t1c", "mr", "snapshot", "t1dixon", "dwi", "ctb"] diff --git a/TPTBox/core/dicom/dicom2nii_utils.py b/TPTBox/core/dicom/dicom2nii_utils.py index 7fb7b65..4adb068 100755 --- a/TPTBox/core/dicom/dicom2nii_utils.py +++ b/TPTBox/core/dicom/dicom2nii_utils.py @@ -289,10 +289,27 @@ def test_name_conflict(json_ob: dict, file: str | Path) -> bool: ``False`` otherwise (file does not exist or content matches). """ if Path(file).exists(): - with open(file, encoding="utf-8") as f: - js = json.load(f) - if "grid" in js: - del js["grid"] + try: + with open(file, encoding="utf-8") as f: + js = json.load(f) + except (UnicodeDecodeError, json.JSONDecodeError, OSError) as e: + # Unreadable JSON at the target path is almost always an artefact of + # a previous crashed / half-written extraction, not an unrelated + # file that the extractor should route around. Delete it and treat + # the slot as free so the caller writes fresh content over it, + # rather than piling up ``_sequ--a`` copies alongside the + # broken original. Fresh-content path relies on the caller passing + # ``override=True`` to `save_json` (the default) or having no file + # at all — both hold here. + Print_Logger().on_warning(f"test_name_conflict: replacing unreadable JSON at {file} ({type(e).__name__}: {e}).") + try: + Path(file).unlink(missing_ok=True) + except OSError as unlink_err: + Print_Logger().on_warning(f"test_name_conflict: could not unlink corrupt JSON {file}: {unlink_err}") + return True # fall back to the "rename around it" path + return False + if "grid" in js: + del js["grid"] return js != json_ob return False diff --git a/TPTBox/core/dicom/dicom_extract.py b/TPTBox/core/dicom/dicom_extract.py index c91c98f..9290b61 100644 --- a/TPTBox/core/dicom/dicom_extract.py +++ b/TPTBox/core/dicom/dicom_extract.py @@ -39,6 +39,13 @@ logger = Print_Logger() +# Modalities/BIDS formats we drop by default. Grayscale Softcopy Presentation +# State (`pr`) DICOMs carry no pixel data — only Referenced SOP Instance UIDs +# and viewer window/level/annotation presets — and always come out as noise +# for automated pipelines. Callers can pass `skip_formats=set()` to keep them, +# or add other formats (e.g. `{"pr", "ko"}` to also drop Key Object Selection). +_DEFAULT_SKIP_FORMATS: set[str] = {"pr", "xa-helper"} + def _next_letter_suffix(s: str, inc: int = 1) -> str: """Increment a letter suffix: a -> b, z -> aa, aa -> ab.""" @@ -57,6 +64,9 @@ def _next_letter_suffix(s: str, inc: int = 1) -> str: return "".join(reversed(result)) +_INC_KEY_MAX_TRIES = 10_000 + + def _inc_key(keys: dict, inc: int = 1, k: str = "sequ", path_exists: Callable[[dict], bool] | None = None) -> None: """Increment the sequence key inside *keys* by appending letter suffixes. @@ -64,6 +74,11 @@ def _inc_key(keys: dict, inc: int = 1, k: str = "sequ", path_exists: Callable[[d until the filename generated from *keys* no longer collides with an existing file on disk. This guarantees the caller never receives keys that would produce a duplicate filename. + + Raises: + RuntimeError: If ``path_exists`` never returns ``False`` within + :data:`_INC_KEY_MAX_TRIES` iterations. Prevents a broken + ``path_exists`` callback from spinning forever. """ def _step() -> None: @@ -88,8 +103,15 @@ def _step() -> None: keys[k] = f"{value}-a" _step() + tries = 0 while path_exists is not None and path_exists(keys): _step() + tries += 1 + if tries >= _INC_KEY_MAX_TRIES: + raise RuntimeError( + f"_inc_key: `path_exists` still True after {tries} increments (current keys[{k!r}]={keys.get(k)!r}). " + "This suggests a mis-configured `path_exists` callback rather than a real filename collision." + ) def _generate_bids_path( @@ -283,6 +305,7 @@ def _export_pdf_from_dicom(dcm_path, out_pdf): def _collect_text(ds, txt_lines: list[str] | None = None): if txt_lines is None: txt_lines = [] + start = len(txt_lines) def _help_collect_text(content_sequence, level: int = 0): for item in content_sequence: @@ -314,6 +337,28 @@ def _help_collect_text(content_sequence, level: int = 0): if hasattr(ds, "ContentSequence"): _help_collect_text(ds.ContentSequence) + + # Fallback for non-SR modalities (Presentation State, Key Object Selection, + # Registration, Fiducials, waveforms, …). These carry no ContentSequence, + # so the SR walker above produces nothing and the previous behaviour left + # an empty .txt. Dump every DICOM element (minus the pixel data blob) — + # the same tag/name/VR/value view a DICOM tag inspector would show. + if len(txt_lines) == start: + try: + iterator = ds.iterall() if hasattr(ds, "iterall") else ds + for elem in iterator: + if getattr(elem, "tag", None) is not None and elem.tag.group == 0x7FE0: + continue # PixelData family + try: + txt_lines.append(str(elem)) + except Exception: # noqa: BLE001 + txt_lines.append(f"({getattr(elem, 'tag', '?')}) ") + except Exception: # noqa: BLE001 + # Last resort: Dataset.__str__ still gives a readable dump. + try: + txt_lines.extend(str(ds).splitlines()) + except Exception: # noqa: BLE001 + pass return txt_lines @@ -346,8 +391,26 @@ def _extract_nii_from_dicom(dicom_out_path, nii_path): ds = dicom_out_path[0] if hasattr(ds, "pixel_array") and len(ds.pixel_array.shape) >= 2: dicom_to_nifti_multiframe(ds, nii_path) - - return True + return True + # Single-DICOM series without usable pixel data (Presentation + # State, Key Object Selection, Registration, Fiducials, some + # RT objects). Previously `return True` here made the caller + # run `_add_grid_info_to_json` on a NIfTI that was never + # written and crash with FileNotFoundError. Also try to lift + # any ContentSequence into a `.txt` sidecar so structured- + # report-style DICOMs don't lose their textual payload; the + # `.json` header sidecar is kept in either case (it already + # holds every non-pixel DICOM tag) so downstream inspection + # still works. Return False so no grid step follows. + logger.on_debug(f"Not exportable (no pixel_array): {Path(nii_path).name}") + try: + txt_path = str(nii_path).replace(".nii.gz", ".txt") + _extract_txt_from_dicom(dicom_out_path, txt_path) + if Path(txt_path).stat().st_size == 0: + Path(txt_path).unlink(missing_ok=True) + except Exception as e: # noqa: BLE001 + logger.on_debug(f"Text dump failed for {Path(nii_path).name}: {e}") + return False except Exception as e: logger.on_debug("Multi-Frame DICOM did not work:", e) ## The PDF dicom lands here @@ -380,8 +443,14 @@ def _extract_nii_from_dicom(dicom_out_path, nii_path): logger.print_error() return False - except Exception: - print(nii_path) + except Exception as e: # noqa: BLE001 + # Any other conversion failure: log and treat as a failed conversion so + # callers don't run downstream steps (e.g. `_add_grid_info_to_json` or + # `_split_multi_echo_dixon`) on a file that was never written. + logger.on_warning(f"_extract_nii_from_dicom: unexpected error on {nii_path}: {type(e).__name__}: {e}") + logger.print_error() + Path(str(nii_path).replace(".nii.gz", ".json")).unlink(missing_ok=True) + return False return True @@ -462,6 +531,7 @@ def _from_dicom_to_nii( skip_localizer: bool = False, parent="rawdata", censor_list=None, + skip_formats: set[str] | None = None, ): """Convert a list of DICOM datasets for one series to a NIfTI file. @@ -479,11 +549,18 @@ def _from_dicom_to_nii( override_subject_name: Optional callable that returns a custom subject name. chunk: Chunk index for multi-stack series; ``None`` triggers automatic splitting. skip_localizer: Skip localizer series when ``True``. + skip_formats: BIDS format labels to drop entirely (no JSON, no sidecar, + no NIfTI). Defaults to :data:`_DEFAULT_SKIP_FORMATS` = ``{"pr"}`` + — Grayscale Softcopy Presentation State DICOMs carry no pixel + data and only reference other series, so they normally contribute + nothing to a downstream pipeline. Pass ``set()`` to keep them. Returns: Path to the generated NIfTI file, ``None`` on failure, or a list of paths when the series was automatically split into multiple stacks. """ + if skip_formats is None: + skip_formats = _DEFAULT_SKIP_FORMATS if censor_list is None: censor_list = [ "StudyDate", @@ -514,6 +591,7 @@ def _from_dicom_to_nii( chunk=i, skip_localizer=skip_localizer, parent=parent, + skip_formats=skip_formats, ) outs.append(o) return outs @@ -540,6 +618,9 @@ def _from_dicom_to_nii( ) if skip_localizer and json_bids.bids_format == "localizer": return + if json_bids.bids_format in skip_formats: + logger.on_debug(f"Skipping {json_bids.bids_format!r} series (in skip_formats): {Path(json_file_name).name}") + return None logger.print(json_file_name, Log_Type.NEUTRAL, verbose=verbose) exist = save_json(simp_json, json_file_name, override=False) # logger.on_debug(exist, Path(nii_path).exists(), nii_path) @@ -556,10 +637,12 @@ def _from_dicom_to_nii( if add_grid: _add_grid_info_to_json(nii_path, json_file_name) - # Multi-echo Philips DIXON (magnitude/phase) arrives as a 4-D NIfTI. - # Split it into per-echo 3-D files with `-eco` appended to `part`. - if json_bids.get("part") in ("magnitude", "phase"): - _split_multi_echo_dixon(Path(nii_path), Path(json_file_name), dcm_data_l) + # Multi-echo Philips DIXON arrives as a 4-D NIfTI. Try to split whenever + # the output is 4-D — `_split_multi_echo_dixon` is a no-op on 3-D input + # and safely returns None. This catches multi-echo series whose `part` + # entity was mapped to something other than "magnitude"/"phase" via + # `dixon_mapping` or `parts_mapping`. + _split_multi_echo_dixon(Path(nii_path), Path(json_file_name), dcm_data_l) return nii_path if add_grid else None @@ -613,7 +696,7 @@ def _split_multi_echo_dixon(nii_path: Path, json_path: Path, dcm_data_l) -> list parent_json = load_json(json_path) if Path(json_path).exists() else {} frames = nii.split_4D_image_to_3D() out_paths: list[Path] = [] - for i, (frame, te) in enumerate(zip(frames, tes)): + for i, (frame, te) in enumerate(zip_strict(frames, tes)): new_nii = _with_echo_suffix(nii_path, i) new_json = _with_echo_suffix(json_path, i) frame.save(new_nii) @@ -686,6 +769,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. @@ -698,36 +807,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 @@ -744,8 +853,13 @@ def _find_all_files(dcm_dirs: Path | list[Path], verbose=False): if verbose: logger.on_neutral("Start file searching") i = 0 - yield dcm_dirs dcm_dirs = dcm_dirs if isinstance(dcm_dirs, list) else [dcm_dirs] + # Yield each root path individually so callers can process a single-directory + # source without descending. Previously the raw `dcm_dirs` was yielded first + # (a list, when a list was passed) and downstream `str(dicom_path)` cast the + # list into a bogus string — no caller could parse it. + for dcm_dir in dcm_dirs: + yield dcm_dir for dcm_dir in dcm_dirs: if dcm_dir.is_dir(): for root, _, files in os.walk(dcm_dir): @@ -779,6 +893,67 @@ def _unzip_files(dicom_zip_path: Path, out_dir: str | Path) -> Path: return dicom_out_path +# Non-DICOM file extensions we can rule out without touching pydicom. Speeds up +# `_read_dicom_files` significantly on trees that mix DICOMs with reports, +# thumbnails, or metadata. Archive extensions are intentionally NOT listed +# here — top-level `.zip` sources are already unpacked by `extract_dicom_folder` +# before `_read_dicom_files` runs, and any residual archive would fail the +# DICM-magic check below. +_NON_DICOM_SUFFIXES = frozenset( + { + ".json", + ".txt", + ".md", + ".pdf", + ".png", + ".jpg", + ".jpeg", + ".tif", + ".tiff", + ".gif", + ".bmp", + ".ini", + ".log", + ".csv", + ".tsv", + ".yaml", + ".yml", + ".html", + ".htm", + ".xml", + ".nii", # already-extracted NIfTI + ".nrrd", + ".mha", + ".mhd", + } +) + + +def _looks_like_dicom(path: Path) -> bool: + """Cheap check whether *path* looks like a DICOM file. + + First rules out common non-DICOM extensions, then reads the first 132 bytes + and checks for the ``DICM`` magic at offset 128 (the standard preamble). + Files without the preamble (deflated/implicit) fall back to a "no extension + and non-empty" heuristic — matches the previous behaviour where every + extensionless file was handed to :func:`pydicom.dcmread`. + """ + suffix = path.suffix.lower() + if suffix in _NON_DICOM_SUFFIXES: + return False + try: + with path.open("rb") as fh: + head = fh.read(132) + except OSError: + return False + if len(head) >= 132 and head[128:132] == b"DICM": + return True + # Some DICOM files skip the 128-byte preamble; retain the old permissive + # behaviour for extensionless files (matches Siemens/Philips exports that + # ship as `IM000001` etc.). + return suffix in ("", ".dcm", ".ima", ".dicom") + + def _read_dicom_files(dicom_out_path: Path) -> tuple[dict[str, list[FileDataset]], dict[str, list[str]]]: """Read DICOM files from a directory and categorize them based on SeriesInstanceUID and type. @@ -794,7 +969,7 @@ def _read_dicom_files(dicom_out_path: Path) -> tuple[dict[str, list[FileDataset] dicom_types: dict[str, list[str]] = {} for _paths in dicom_out_path.rglob("*"): path = Path(_paths) - if path.is_file(): + if path.is_file() and _looks_like_dicom(path): try: dcm_data = pydicom.dcmread(path, defer_size="1 KB", force=True) # , stop_before_pixels=True try: @@ -827,6 +1002,37 @@ def _read_dicom_files(dicom_out_path: Path) -> tuple[dict[str, list[FileDataset] return dicom_files, _filter_file_type(dicom_types) +def _split_by_echo_numbers(dicoms: list[FileDataset]) -> tuple[list[FileDataset], dict[int, list[FileDataset]]]: + """Split DICOMs into a single-echo subset (for stack detection) and per-echo buckets. + + Multi-echo Philips series (e.g. Philips mDIX quant) place M echoes at each + of N slice positions inside a single ImageType sub-group. Consecutive DICOMs + sorted by InstanceNumber then sit at the same ``ImagePositionPatient``, which + makes the direction-based stack detection in :func:`_classic_get_grouped_dicoms` + produce ``0 / 0 = NaN`` and split the stack into random chunks. + + Returns a ``(single_echo_subset, per_echo_buckets)`` pair. ``single_echo_subset`` + is the DICOMs of the smallest ``EchoNumbers`` value (i.e. one DICOM per slice + position) and can be handed straight to the stack detector. ``per_echo_buckets`` + maps each ``EchoNumbers`` value to its DICOMs, so the caller can put the other + echoes back after stack detection. + + No-op when there is 0 or 1 distinct ``EchoNumbers`` — the full input is returned + as ``single_echo_subset`` with an empty ``per_echo_buckets``. + """ + per_echo: dict[int, list[FileDataset]] = {} + for d in dicoms: + try: + en = int(getattr(d, "EchoNumbers", 0) or 0) + except (TypeError, ValueError): + en = 0 + per_echo.setdefault(en, []).append(d) + if len([k for k in per_echo if k > 0]) <= 1: + return list(dicoms), {} + keep = min(k for k in per_echo if k > 0) + return list(per_echo[keep]), per_echo + + def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[FileDataset]]: """Group DICOM slices into spatially contiguous stacks by analysing slice direction. @@ -835,6 +1041,12 @@ def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[Fil spatial acquisition stack. Groups with three or fewer slices are collected into a single catch-all group at the end. + For multi-echo series (multiple ``EchoNumbers`` values), the stack detection + runs on a single-echo subset — DICOMs at the same ``ImagePositionPatient`` + would otherwise produce ``0 / 0 = NaN`` in the direction test and split one + real acquisition into random stacks. The remaining echoes are re-attached + to the returned groups by ``ImagePositionPatient``. + Args: dicom_input: Flat list of pydicom ``FileDataset`` objects for a single series. @@ -842,8 +1054,10 @@ def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[Fil List of groups, where each group is a list of ``FileDataset`` objects belonging to the same spatial stack. """ + detection_set, per_echo = _split_by_echo_numbers(dicom_input) + # Order all dicom files by InstanceNumber - dicoms = sorted(dicom_input, key=lambda x: x.InstanceNumber) + dicoms = sorted(detection_set, key=lambda x: x.InstanceNumber) # now group per stack grouped_dicoms: list[list[FileDataset]] = [[]] # list with first element a list @@ -857,8 +1071,14 @@ def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[Fil current_direction = None # if the stack number decreases we moved to the next stack if previous_position is not None: - current_direction = np.array(dicom_.get("ImagePositionPatient", 0)) - previous_position - current_direction = current_direction / np.linalg.norm(current_direction) + delta = np.array(dicom_.get("ImagePositionPatient", 0)) - previous_position + norm = float(np.linalg.norm(delta)) + # Zero-length delta = two DICOMs at the same position (residual + # multi-echo where the pre-filter didn't remove all duplicates, + # or a genuine repeated slice). Skip the direction update so we + # don't emit NaN and split the stack. + if norm > 1e-6: + current_direction = delta / norm if ( current_direction is not None @@ -870,7 +1090,8 @@ def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[Fil stack_index += 1 else: previous_position = np.array(dicom_.get("ImagePositionPatient", 0)) - previous_direction = current_direction + if current_direction is not None: + previous_direction = current_direction if stack_index >= len(grouped_dicoms): grouped_dicoms.append([]) @@ -884,6 +1105,33 @@ def _classic_get_grouped_dicoms(dicom_input: list[FileDataset]) -> list[list[Fil out.append(i) if len(others) != 0: out.append(others) + + # Re-attach the other echoes to whichever stack their spatial position + # belongs to. Only meaningful when the input is genuinely multi-echo AND + # the stack detector actually split into more than one group. + if per_echo and len(out) > 1: + + def _key(d: FileDataset) -> tuple: + return tuple(float(v) for v in d.get("ImagePositionPatient", (0.0, 0.0, 0.0))) + + pos_to_stack: dict[tuple, int] = {} + for i, group in enumerate(out): + for d in group: + pos_to_stack[_key(d)] = i + detection_keep = min(k for k in per_echo if k > 0) + for en, echo_dicoms in per_echo.items(): + if en == detection_keep: + continue # already placed via the detection set + for d in echo_dicoms: + idx = pos_to_stack.get(_key(d)) + if idx is not None: + out[idx].append(d) + elif out: + # Unknown position: dump into the catch-all "others" tail + out[-1].append(d) + elif per_echo and len(out) == 1: + # Single-stack multi-echo: return every echo in one group. + out = [list(dicom_input)] return out @@ -949,6 +1197,7 @@ def extract_dicom_folder( censor_list: list | None = None, skip_already_extracted: bool = True, force_rescan: bool = False, + skip_formats: set[str] | None = None, ) -> dict: """Extract DICOM files from a directory or list of directories, convert them to NIfTI format, and store the output. @@ -966,7 +1215,13 @@ def extract_dicom_folder( validate_orientation (bool, optional): Enable ``dicom2nifti`` orientation validation. Defaults to True. validate_orthogonal (bool, optional): Enable ``dicom2nifti`` orthogonality validation. Defaults to False. validate_slice_increment (bool, optional): Enable ``dicom2nifti`` slice-increment validation. Defaults to True. - n_cpu (int, optional): Number of CPU cores to use for parallel processing. Defaults to 1 (sequential). + n_cpu (int | None, optional): Threading policy for per-series conversion. + ``1`` (default) processes series sequentially. ``>1`` uses that many + worker threads. ``None`` hands ``max_workers=None`` to + :class:`concurrent.futures.ThreadPoolExecutor`, whose Python default + is ``min(32, os.cpu_count() + 4)`` — effectively "use all cores, up + to 32". Note that DICOM extraction is I/O-heavy, so threads (not + processes) tend to be the right knob. override_subject_name (Callable[[dict, Path], str] | None, optional): Callable receiving the parsed DICOM header dict and file path; returns the subject id to use in the BIDS output. Defaults to None. skip_localizer (bool, optional): If True, skip series identified as scanner localisers. Defaults to True. @@ -980,6 +1235,12 @@ def extract_dicom_folder( added subjects. Defaults to True. force_rescan (bool, optional): If True, bypass the fast-skip marker and re-read every DICOM. Defaults to False. + skip_formats (set[str] | None, optional): BIDS format labels to drop + entirely (no JSON, no sidecar, no NIfTI). Defaults to + :data:`_DEFAULT_SKIP_FORMATS` = ``{"pr"}`` — Grayscale Softcopy + Presentation State DICOMs carry no pixel data and only reference + other series, so they normally contribute nothing to a downstream + pipeline. Pass ``set()`` to keep them. Returns: dict: A dictionary with keys representing DICOM series and values as paths to the generated NIfTI files. @@ -1001,18 +1262,16 @@ 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. - 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)) - ): + # 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 _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"): @@ -1039,6 +1298,7 @@ def process_series(key, files, parts): skip_localizer=skip_localizer, parent=parent, censor_list=censor_list, + skip_formats=skip_formats, ) # Process in parallel or sequentially based on n_cpu @@ -1061,14 +1321,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 54db4da..9cc8a84 100644 --- a/TPTBox/core/dicom/dicom_header_to_keys.py +++ b/TPTBox/core/dicom/dicom_header_to_keys.py @@ -68,6 +68,42 @@ ".*mp?ra?ge?.*": "MPR", ".*mip.*": "MIP", "b0map": "b0map", + # Specific quantitative / specialised patterns MUST come before the greedy + # ``.*t2.*`` / ``.*t1.*`` catch-alls below — otherwise "T2 STAR" / "T1 MAP" + # get misclassified as plain T2w / T1w on the first-match win. + ".*mp2rage.*": "MP2RAG", + ".*t2\\s*star.*": "T2star", + r".*t2\*.*": "T2star", + ".*r2\\s*star.*": "R2star", + r".*r2\*.*": "R2star", + ".*swi.*": "SWI", + ".*t1\\s*map.*": "T1map", + ".*t2star\\s*map.*": "T2starmap", + ".*t2\\s*map.*": "T2map", + # Multi-echo VIBE / DIXON — NAKO Siemens ``ME_vibe_fatquant_*`` and the + # generic ``mevibe``/``fatquant``/``fatfrac``/``pdff`` labels. Placed before + # ``.*mdix.*`` so the multi-echo classification wins where both apply. + ".*me[_\\s]?vibe.*": "mevibe", + ".*mevibe.*": "mevibe", + ".*fatquant.*": "mevibe", + ".*fatfrac.*": "dixon", + ".*pdff.*": "dixon", + ".*ideal.*": "dixon", # GE's Dixon variant + # Philips-specific localizers / reference scans that the existing "pilot" + # / "scout" entries above don't catch. + ".*survey.*": "localizer", + ".*ref\\s*scan.*": "localizer", + ".*smartexam.*": "localizer", + # Philips ``3DI_MC_HR`` = 3D-Inflow Motion-Corrected High-Resolution + # (TOF-MRA MIP projections; ImageType carries ``PROJECTION IMAGE``). + # Placed before the generic angio/tof rules below because the raw + # ``3di_mc_hr`` string contains neither "tof" nor "angio". + ".*3di[_-]?mc.*": "TOF", + # Post-contrast dynamic MR (KM = Kontrastmittel). Requires all three + # tokens so a generic ``dyn`` doesn't over-match. + ".*dyn.*post.*km.*": "DCE", + # German ``Halsgefäße`` neck-vessel angiography. + ".*halsgef.*": "angio", ".*t2.*": "T2w", ".*t1.*": "T1w", ".*dixon.*": "dixon", @@ -82,6 +118,7 @@ ".*sub.*": "subtraction", ".*dynamik.*": "DCE", ".*mdix.*": "dixon", + ".*mdixon.*": "dixon", ".*s3d.*": "s3D", ".*flip37.*": "s3D", ".*trak.*": "PWI", @@ -104,6 +141,85 @@ } +def _single_echo_for_plane(dicoms: list[pydicom.FileDataset]) -> list[pydicom.FileDataset]: + """Return a DICOM subset with one echo per slice position for plane detection. + + Multi-echo Philips DIXON (e.g. "mDIX quant") exports N slice positions × M + echos into a single DICOM sub-group. The M copies at each `ImagePositionPatient` + collapse the slice axis to ≈0 in `dicom2nifti.common.create_affine`, so after + clamping by `hires_threshold` every zoom is ~1 and the series is misdetected + as isotropic. Keep only the smallest `EchoNumbers` value so each spatial + position is represented once. No-op when the tag is absent or constant. + """ + en_values = set() + for d in dicoms: + try: + en = int(getattr(d, "EchoNumbers", 0) or 0) + except (TypeError, ValueError): + continue + if en > 0: + en_values.add(en) + if len(en_values) <= 1: + return dicoms + keep = min(en_values) + return [d for d in dicoms if int(getattr(d, "EchoNumbers", 0) or 0) == keep] + + +def _apply_view_keys(keys: dict, get: Callable) -> None: + """Populate `acq` / `part` from DICOM ViewPosition + Laterality tags. + + Used by the 2D-modality fallback in :func:`extract_keys_from_json`. Without + this, MG series with four views (R-CC, L-CC, R-MLO, L-MLO), radiographs + with AP / PA / LAT projections, and ophthalmic photos of both eyes would + all collapse onto the same BIDS filename and clobber each other. + + Convention chosen (matching this codebase's flexible ``acq`` usage): + + * ``ViewPosition`` (e.g. ``CC``, ``MLO``, ``AP``, ``PA``, ``LAT``) → + lowercased into ``acq``. Only overwrites the existing ``acq`` when the + plane-detector returned ``None`` or ``"iso"``, both of which are + meaningless for single-slice imagery. + * ``ImageLaterality`` / ``Laterality`` (``L`` / ``R``, or ophthalmic + ``OS`` / ``OD`` → mapped to ``L`` / ``R``) → ``part``, only when + ``part`` is not already set by the DIXON / ImageType branches above. + """ + view = get("ViewPosition") + if view: + view_clean = str(view).lower().strip("-.") + if view_clean and keys.get("acq") in (None, "iso"): + keys["acq"] = view_clean + laterality = get("ImageLaterality") or get("Laterality") + if laterality: + lat = str(laterality).upper() + # Ophthalmic (OS = oculus sinister = left, OD = oculus dexter = right) + # normalises to the same L/R vocabulary that radiography uses. + lat = {"OS": "L", "OD": "R"}.get(lat, lat) + if lat in {"L", "R", "B"} and keys.get("part") is None: + keys["part"] = lat.lower() + + +def _apply_bodypart_key(keys: dict, get: Callable) -> None: + """Populate ``desc`` from DICOM ``BodyPartExamined`` when nothing else has set it. + + ``BodyPartExamined`` (0018,0015) is a semi-standardised free-text tag with + common values like ``ABDOMEN``, ``PELVIS``, ``ABDOMENPELVIS``, ``CHEST``, + ``HEAD``, ``NECK``, ``SPINE``, ``KNEE``, ``HIP``, ``BREAST``. It is the + main discriminator when a single session contains scans of several body + regions and the ``SeriesDescription`` is not informative — typical for + plain radiography, ultrasound, nuclear medicine, and RT objects. Only + written when ``keys['desc']`` is empty, so an earlier branch that already + assigned ``desc`` (e.g. the SR / report path) wins. + """ + if keys.get("desc") is not None: + return + body = get("BodyPartExamined") + if not body: + return + val = str(body).lower().strip("-.") + if val: + keys["desc"] = val + + def get_plane_dicom(dicoms: list[pydicom.FileDataset] | NII, hires_threshold: float = 0.8) -> str | None: """Determine the acquisition plane from a DICOM series or NIfTI image. @@ -126,7 +242,7 @@ def get_plane_dicom(dicoms: list[pydicom.FileDataset] | NII, hires_threshold: fl if isinstance(dicoms, NII): return dicoms.get_plane(res_threshold=hires_threshold) try: - sorted_dicoms = common.sort_dicoms(dicoms) + sorted_dicoms = common.sort_dicoms(_single_echo_for_plane(dicoms)) affine, _ = common.create_affine(sorted_dicoms) plane_dict = {"S": "ax", "I": "ax", "L": "sag", "R": "sag", "A": "cor", "P": "cor"} axc = np.array(nio.aff2axcodes(affine)) @@ -149,7 +265,23 @@ def get_plane_dicom(dicoms: list[pydicom.FileDataset] | NII, hires_threshold: fl else: plane = "iso" return plane # noqa: TRY300 - except Exception: + except (AttributeError, IndexError, KeyError, TypeError): + # Not usable image geometry: non-imaging DICOMs legally lack + # `ImagePositionPatient` / `ImageOrientationPatient` (AttributeError), + # empty lists trip `create_affine` on `dicoms[0]` (IndexError), and + # callers that hand in dicts or other pydicom-shaped-but-not-really + # objects raise KeyError / TypeError. All of these mean "no plane to + # compute" — return None silently instead of surfacing the noise. + return None + except Exception as e: # noqa: BLE001 + # Log so a downstream `acq-None` filename can be traced back to its cause, + # instead of the plane-detection silently swallowing every failure. + try: + from TPTBox import Print_Logger + + Print_Logger().on_warning(f"get_plane_dicom: plane detection failed ({type(e).__name__}: {e}); returning None.") + except Exception: # noqa: BLE001 + pass return None @@ -216,18 +348,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 +384,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: @@ -270,9 +427,9 @@ def _get(key, default=None): if session: keys["ses"] = _get("StudyDate", keys.get("ses")) if isinstance(dcm_data_l, (str, Path, NII)): - keys["acq"] = to_nii(dcm_data_l).get_plane(1) + keys["acq"] = to_nii(dcm_data_l).get_plane(0.8) else: - keys["acq"] = get_plane_dicom(dcm_data_l, 1) + keys["acq"] = get_plane_dicom(dcm_data_l, 0.8) keys["part"] = dixon_mapping.get(_get("ProtocolName", "NO-PART").split("_")[-1]) sequ = _get("SeriesNumber", None) @@ -317,9 +474,19 @@ def _get(key, default=None): found = False if modality == "ct": mri_format = "ct" + _apply_bodypart_key(keys, _get) elif modality.lower() == "pt": mri_format = "pet" + _apply_bodypart_key(keys, _get) elif modality == "xa": # Angiography + # Helper / non-imaging XA payload: SECONDARY captures, referenced- + # image thumbnails, and exam-protocol screenshots have modality XA + # but no diagnostic pixel volume. Route to `xa-helper` so + # `_DEFAULT_SKIP_FORMATS` drops them at the caller instead of + # cluttering the output with unusable derivations. + if any(t in image_type for t in ("SECONDARY", "REFIMAGE", "EXAM PROTOCOL")): + mri_format = "xa-helper" + return mri_format, keys, ".nii.gz" biplane = False if "BIPLANE A" in image_type or "SINGLE A" in image_type: keys["acq"] = "A" @@ -349,6 +516,17 @@ def _get(key, default=None): mri_format = "DSA" else: mri_format = "XA" + # Manufacturer-agnostic ImageType fallbacks — override the plain + # `XA` fallback with more specific labels when the header + # unambiguously says so. `ORIGINAL` runs are live fluoro; a + # `DERIVED PRIMARY` subtracted plane is a DSA. Keeps the + # SeriesDescription-string checks above as first-pass and only + # kicks in when they didn't resolve past `XA`. + if mri_format == "XA": + if "ORIGINAL" in image_type: + mri_format = "fluroscopy" + elif "DERIVED" in image_type and "PRIMARY" in image_type: + mri_format = "DSA" elif modality == "mr": for key, mri_format_new in map_series_description_to_file_format.items(): regex = re.compile(key) @@ -370,13 +548,104 @@ def _get(key, default=None): " km " in series_description.lower() or series_description.startswith("km") or series_description.endswith("km") ) and keys.get("ce") is None: keys["ce"] = "ContrastAgent" + # 2D projections (Philips MIP views of a TOF-MRA source volume, + # or any DICOM tagged with `PROJECTION IMAGE` in ImageType). Same + # SeriesDescription as the 3D source volume, but the pixel data + # is a MIP. Route these to the `MIP` format so the 3D recons keep + # `TOF` / `angio` labels and the MIPs live under `sub-*/ses-*/MIP/`. + if any("PROJECTION" in str(t).upper() for t in image_type): + mri_format = "MIP" elif modality.lower() == "pdf": return "report", keys, ".pdf" elif modality.lower() == "sr": keys["desc"] = _get("SeriesDescription", None) return "report", keys, ".txt" + # Non-imaging metadata DICOMs: Presentation State, Key Object Selection, + # Registration, Fiducials, Real World Value Map, Plan, Slide Stainer. + # Also physiological waveforms (RESP, HD, ECG, EPS) and ophthalmic + # measurements (AR, KER, LEN, VA, OPV, OPM) — none of these carry a + # NIfTI-shaped pixel volume. Route them through the same `.txt` report + # path as SR so the caller neither writes an empty NIfTI nor crashes + # in `_add_grid_info_to_json` on a file that was never produced. + elif modality.lower() in { + "pr", + "ko", + "reg", + "fid", + "rwv", + "plan", + "stain", + "resp", + "hd", + "ecg", + "eps", + "ar", + "ker", + "len", + "va", + "opv", + "opm", + }: + keys["desc"] = _get("SeriesDescription", None) + return modality.lower(), keys, ".txt" + # Sensible defaults for the remaining common imaging modalities so we can + # keep converting instead of raising on every non-CT/PET/MR/XA series. + # Format names mirror BIDS conventions where they exist and fall back to + # the lowercased DICOM modality tag otherwise (e.g. `us`, `nm`, `sc`). + # For 2D modalities we also lift the DICOM ViewPosition / Laterality tags + # into `acq`, otherwise files that only differ by view (R-CC vs L-CC vs + # R-MLO vs L-MLO for MG, AP vs PA vs LAT for DX/CR) would all collapse + # to the same BIDS name. + elif modality.lower() in {"cr", "dx", "rg", "px", "io", "mg"}: + # 2D X-ray family: computed / digital radiography, general radiographic, + # panoramic, intra-oral, mammography. Kept under one `xray` bucket. + mri_format = "xray" + _apply_view_keys(keys, _get) + _apply_bodypart_key(keys, _get) + elif modality.lower() == "us": + mri_format = "us" # ultrasound + _apply_view_keys(keys, _get) + _apply_bodypart_key(keys, _get) + elif modality.lower() == "nm": + mri_format = "nm" # nuclear medicine (planar/SPECT) + # Radiopharmaceutical (tracer) is the useful discriminator for NM — + # e.g. FDG, PSMA, DOTATATE. When present, surface it as `ce`. + tracer = _get("Radiopharmaceutical") + if tracer and keys.get("ce") is None: + keys["ce"] = tracer + _apply_bodypart_key(keys, _get) + elif modality.lower() == "sc": + mri_format = "sc" # secondary capture (screenshots, derived stills) + _apply_bodypart_key(keys, _get) + elif modality.lower() in {"op", "xc"}: + mri_format = "photo" # ophthalmic / external photography + # Ophthalmic photos: OS = left eye, OD = right eye → same L/R signal + # as radiography Laterality; reuse the same helper. + _apply_view_keys(keys, _get) + _apply_bodypart_key(keys, _get) + elif modality.lower() == "es": + mri_format = "endoscopy" + _apply_bodypart_key(keys, _get) + elif modality.lower() in {"rtimage", "rtstruct", "rtdose", "rtplan"}: + mri_format = modality.lower() # radiotherapy objects + _apply_bodypart_key(keys, _get) + elif modality.lower() == "ot": + mri_format = "ot" # explicit "Other" modality + _apply_bodypart_key(keys, _get) else: - raise NotImplementedError(f"modality='{modality}', ({modalities.get(modality.upper(), 'Non Standard Modality key')})") + # Unknown modality — warn once and fall back to a mri_format derived + # from the modality tag so extraction can still complete. Callers + # that really need to reject unknown modalities can inspect the + # returned mri_format. + from TPTBox import Print_Logger + + Print_Logger().on_warning( + f"extract_keys_from_json: unhandled modality={modality!r} " + f"({modalities.get(modality.upper(), 'Non Standard Modality key')}); " + "falling back to modality tag as mri_format." + ) + mri_format = str(modality).lower() or "mr" + _apply_bodypart_key(keys, _get) # ".*sub.*t1.*": "subtraktion", # "subtraktion.*t1.*": "subtraktion", diff --git a/TPTBox/core/dicom/dicom_renamer.py b/TPTBox/core/dicom/dicom_renamer.py new file mode 100644 index 0000000..78f5708 --- /dev/null +++ b/TPTBox/core/dicom/dicom_renamer.py @@ -0,0 +1,735 @@ +"""Post-hoc BIDS renamer for datasets already extracted from DICOM. + +Walk a BIDS-style dataset root, re-run +:func:`~TPTBox.core.dicom.dicom_header_to_keys.extract_keys_from_json` on +every sidecar JSON, and rename each file family (`nii.gz` + `json` + `.txt` + +…) to the new BIDS path the current key logic produces. Two use cases: + +1. **Sanitising illegal subject ids.** BIDS forbids ``_`` inside an entity + value (it is the entity separator). Any subject imported with an id like + ``180217_375491`` produced ``sub-180217_375491_ses-...`` filenames that + downstream BIDS parsers reject or misinterpret. Leading ``-`` after + ``sub-`` is the same category of problem. Sanitisation strips those. +2. **Compact numeric ids.** Long random-looking subject ids (``sub-0bCAY6ARpDo``) + are unhandy; passing ``subject_prefix="ID"`` rewrites them to + ``sub-ID001``, ``sub-ID002``, … The mapping is written to + ``//subject_map.tsv`` so the original ids stay + recoverable. + +Additionally, any change picked up upstream in ``extract_keys_from_json`` +(e.g. modality-fallback improvements, view/laterality/BodyPartExamined +extraction) becomes effective on-disk without a full DICOM re-extraction — +the JSON sidecar carries the same header dict the DICOM did. +""" + +from __future__ import annotations + +import csv +import re +from collections.abc import Iterable +from pathlib import Path + +import numpy as np + +from TPTBox import BIDS_FILE, Print_Logger +from TPTBox.core.dicom.dicom2nii_utils import load_json +from TPTBox.core.dicom.dicom_header_to_keys import extract_keys_from_json + + +class _FakeDicom: + """Minimal stand-in for a pydicom Dataset. + + Supplies just the attributes ``extract_keys_from_json`` reads before the + plane branch. + """ + + def __init__(self, filename: str) -> None: + self.filename = filename + + +class _FakeDicomList: + """Tiny list-shaped wrapper for :func:`extract_keys_from_json`. + + Steers upstream away from loading the NIfTI just to redo plane detection. + Supports the two accesses upstream cares about: + * ``lst[0].filename`` — used by the override-subject-name path. + * Iteration — used by ``get_plane_dicom``; we yield nothing so it exits + via the empty-affine `IndexError`, which our expanded silent-catch in + ``get_plane_dicom`` turns into ``None`` without noise. + """ + + def __init__(self, json_path: Path) -> None: + self._stub = _FakeDicom(str(json_path)) + + def __getitem__(self, _i: int) -> _FakeDicom: + return self._stub + + def __iter__(self): + return iter(()) + + def __len__(self) -> int: + return 0 + + +def _plane_from_grid(grid: dict, hires_threshold: float = 0.8) -> str | None: + """Compute the acquisition-plane label from a sidecar ``grid`` dict. + + Mirrors :func:`~TPTBox.core.dicom.dicom_header_to_keys.get_plane_dicom`'s + core zoom/axcodes logic but reads ``spacing`` + ``orientation`` directly + from the sidecar (populated by ``_add_grid_info_to_json`` during + extraction) instead of re-loading the NIfTI. Lets the renamer stay + header-only and process a full dataset in seconds instead of hours. + """ + try: + zooms = np.asarray(grid["spacing"], dtype=float) + orient = grid["orientation"] + except (KeyError, TypeError): + return None + if zooms.size < 3 or not isinstance(orient, (list, tuple)) or len(orient) < 3: + return None + zooms = np.where(zooms == 0, 1.0, zooms) + if hires_threshold is not None: + zooms = np.maximum(zooms, hires_threshold) + zms = np.around(zooms, 1) + plane_dict = {"S": "ax", "I": "ax", "L": "sag", "R": "sag", "A": "cor", "P": "cor"} + ix_max = zms == np.amax(zms) + num_max = int(np.count_nonzero(ix_max)) + axc = np.array(orient[:3]) + if num_max == 2: + return plane_dict.get(axc[~ix_max][0]) + if num_max == 1: + return plane_dict.get(axc[ix_max][0]) + return "iso" + + +logger = Print_Logger() + + +# BIDS entity value grammar allows [A-Za-z0-9] plus limited punctuation. +# `_` is the entity separator and is outright forbidden inside a value. +# A leading `-` inside a value is not spec-forbidden but the parser treats +# it as a new empty-value entity, which corrupts the split. +_ILLEGAL_IN_VALUE = re.compile(r"[_]+") + + +def _sanitize_entity_value(raw: str) -> str: + """Return *raw* with characters that break BIDS entity parsing removed. + + ``_`` (entity separator) is dropped, and a leading ``-`` (which would look + like an empty preceding entity) is stripped. Empty results become + ``"unnamed"`` so callers never build ``sub-_ses-...``-style filenames. + """ + clean = _ILLEGAL_IN_VALUE.sub("", raw).lstrip("-") + return clean or "unnamed" + + +def _iter_json_sidecars(root: Path) -> Iterable[Path]: + """Yield every sidecar JSON under *root*, skipping cache / hidden dirs.""" + for p in sorted(root.rglob("*.json")): + rel = p.relative_to(root) + if any(part.startswith(".") for part in rel.parts): + continue + # Skip our own translation file if it happens to live under root. + if p.name == "subject_map.tsv": + continue + yield p + + +def _list_subject_folders(root: Path) -> list[Path]: + """Return every ``sub-*`` folder directly under *root*, sorted.""" + return sorted(p for p in root.iterdir() if p.is_dir() and p.name.startswith("sub-")) + + +def _current_sub_id(folder: Path) -> str: + """Strip the leading ``sub-`` prefix from a subject folder name.""" + return folder.name[len("sub-") :] + + +def _read_existing_subject_map(dataset_root: Path, info_dir: str) -> dict[str, str]: + """Load a previously written ``subject_map.tsv`` if one exists. + + Returns ``old_sub -> new_sub`` from every row (session columns are + ignored — subject-level identity is the only thing we need to keep + re-runs stable). Missing file → empty dict. + """ + path = dataset_root / info_dir / "subject_map.tsv" + if not path.is_file(): + return {} + saved: dict[str, str] = {} + with path.open(encoding="utf-8") as fh: + reader = csv.DictReader(fh, delimiter="\t") + for row in reader: + old = row.get("old_sub") + new = row.get("new_sub") + if old and new: + saved[old] = new + return saved + + +def _build_subject_map( + subject_folders: list[Path], + subject_prefix: str | None, + subject_number_width: int, + existing_map: dict[str, str] | None = None, +) -> dict[str, str]: + """Map every existing subject id to its new BIDS-legal id. + + * ``subject_prefix=None`` — sanitise in place (drop ``_``, leading ``-``). + * ``subject_prefix="ID"`` — reassign each subject to ``ID001``, ``ID002``, + … in the natural sort order of the folder listing. Prefix itself is + sanitised the same way so the caller can't accidentally reintroduce + ``_`` via the prefix. + + ``existing_map`` (loaded from ``/subject_map.tsv``) makes + re-runs safe. Folders whose id is already a value in ``existing_map`` + (i.e. they were renamed on a previous pass) map to themselves; folders + whose id is still a key get their previously assigned new id back; + only truly fresh subjects get a new number, taken from the next slot + beyond the highest already-used one. Same semantics apply to the + sanitise-in-place branch — an already-sanitised id stays as-is. + """ + existing_map = existing_map or {} + already_new: set[str] = set(existing_map.values()) + # Seed the output with every historic entry so the on-disk translation + # table keeps growing rather than shrinking. Folders whose old-name + # already moved away are still in the map as ``old_random_id -> new_id`` + # rows from the previous run; without this seeding they'd be dropped + # from the file on the next write and the audit trail would be lost. + mapping: dict[str, str] = dict(existing_map) + if subject_prefix is None: + for folder in subject_folders: + old = _current_sub_id(folder) + if old in existing_map: + mapping[old] = existing_map[old] + elif old in already_new: + mapping[old] = old + else: + mapping[old] = _sanitize_entity_value(old) + return mapping + clean_prefix = _sanitize_entity_value(subject_prefix) + used_numbers: set[int] = set() + id_re = re.compile(rf"^{re.escape(clean_prefix)}(\d+)$") + for new in already_new: + m = id_re.match(new) + if m: + used_numbers.add(int(m.group(1))) + for folder in subject_folders: + old = _current_sub_id(folder) + if old in existing_map: + mapping[old] = existing_map[old] + m = id_re.match(existing_map[old]) + if m: + used_numbers.add(int(m.group(1))) + continue + m = id_re.match(old) + if m: # folder is already numbered — keep it as identity. + mapping[old] = old + used_numbers.add(int(m.group(1))) + next_free = (max(used_numbers) + 1) if used_numbers else 1 + for folder in subject_folders: + old = _current_sub_id(folder) + if old in mapping: + continue + while next_free in used_numbers: + next_free += 1 + mapping[old] = f"{clean_prefix}{next_free:0{subject_number_width}d}" + used_numbers.add(next_free) + next_free += 1 + return mapping + + +def _write_subject_map( + dataset_root: Path, + info_dir: str, + mapping: dict[str, str], + session_mapping: dict[tuple[str, str], str] | None = None, +) -> Path: + """Persist ``old_sub -> new_sub`` (plus optional session mapping) as TSV. + + TSV columns: ``old_sub``, ``new_sub``, ``old_ses``, ``new_ses``. When no + session sanitisation happened the ses columns are left empty for the + subject-level row and one row per subject is written; otherwise there is + one row per (subject, session) pair. Written to + ``//subject_map.tsv``; existing content is + overwritten so re-running the renamer keeps the file authoritative. + """ + out_dir = dataset_root / info_dir + out_dir.mkdir(parents=True, exist_ok=True) + path = out_dir / "subject_map.tsv" + with path.open("w", newline="", encoding="utf-8") as fh: + writer = csv.writer(fh, delimiter="\t") + writer.writerow(["old_sub", "new_sub", "old_ses", "new_ses"]) + seen_pairs: set[tuple[str, str]] = set() + if session_mapping: + for (old_sub, old_ses), new_ses in sorted(session_mapping.items()): + new_sub = mapping.get(old_sub, old_sub) + writer.writerow([old_sub, new_sub, old_ses, new_ses]) + seen_pairs.add((old_sub, old_ses)) + for old_sub, new_sub in sorted(mapping.items()): + if any(pair[0] == old_sub for pair in seen_pairs): + continue + writer.writerow([old_sub, new_sub, "", ""]) + return path + + +def _new_bids_path_for( + json_path: Path, + dataset_root: Path, + parent: str, + new_sub_id: str, + session: bool, + make_subject_chunks: int, +) -> tuple[Path, dict]: + """Re-derive the target BIDS path from an existing sidecar JSON. + + The heavy lifting is delegated to :func:`extract_keys_from_json` + + :func:`_generate_bids_path`, exactly as + :func:`~TPTBox.core.dicom.dicom_extract._from_dicom_to_nii` uses them at + extraction time. Two callback tweaks let us reuse the extract logic here: + + * ``override_subject_name`` returns *new_sub_id* verbatim, bypassing the + DICOM-header derivation of the subject id. + * ``dcm_data_l`` is set to the NIfTI path on disk so + :func:`get_plane_dicom`'s ``to_nii`` branch runs — no DICOM headers are + touched. + + Returns ``(new_json_path, new_keys)``; the caller applies the rename with + :meth:`BIDS_FILE.rename_files`. + """ + raw_json = load_json(json_path) + grid = raw_json.get("grid") + simp_json = {k: v for k, v in raw_json.items() if k != "grid"} + # Pre-compute the plane from the sidecar `grid` block; steer + # `extract_keys_from_json` away from re-loading the NIfTI just to redo + # what the sidecar already recorded. We hand a tiny stub list in as + # `dcm_data_l` — it responds to the two accesses upstream cares about + # (``dcm_data_l[0].filename`` for the override-subject-name callback and + # iteration for the plane detector), and the plane detector then returns + # None silently under our expanded exception handler. + pre_plane = _plane_from_grid(grid) if isinstance(grid, dict) else None + dcm_stub = _FakeDicomList(json_path) + (mri_format, keys, _ending) = extract_keys_from_json( + simp_json, + dcm_stub, # type: ignore[arg-type] + session=session, + override_subject_name=lambda _sj, _p: new_sub_id, + ) + if pre_plane is not None: + keys["acq"] = pre_plane + # Build the target path ourselves. `_generate_bids_path` composes via + # `BIDS_FILE.get_changed_bids` which decomposes the current path to + # infer the folder layout — on flat datasets (parent="") that swallows + # the `sub-*` folder into the "parent" slot and drops it when we then + # override parent. Composing `path=sub-/ses-` ourselves keeps + # the layout regardless of how the source dataset happens to be laid + # out on disk. + sub = keys.get("sub") or new_sub_id + ses = keys.get("ses") + sub_folder = f"sub-{sub}" + if make_subject_chunks: + sub_folder = f"{sub[:make_subject_chunks]}/{sub_folder}" + path = f"{sub_folder}/ses-{ses}" if ses else sub_folder + src = BIDS_FILE(json_path, dataset_root, verbose=False) + new_bids = src.get_changed_bids( + file_type="json", + parent=parent, + path=path, + additional_folder=mri_format, + bids_format=mri_format, + make_parent=False, + info=keys, + non_strict_mode=True, + ) + return Path(new_bids.file["json"]), keys + + +def _rename_family(json_path: Path, new_json_path: Path, dataset_root: Path, dry_run: bool) -> list[tuple[Path, Path]]: + """Move every sibling of *json_path* to the *new_json_path* stem. + + Extension detection reuses :meth:`BIDS_FILE.file`, which collects + ``.json`` / ``.nii.gz`` / ``.mrk.json`` etc. as one BIDS entry. + In addition we glob the JSON's folder for any file sharing the exact + base stem — ``BIDS_FILE.file`` currently misses DWI companion files + (``.bval`` / ``.bvec``) and other non-standard extensions that live + next to the primary NIfTI. Same-stem globbing is safe because BIDS + guarantees only truly-paired sidecars share the stem. + """ + if json_path == new_json_path: + return [] + bf_current = BIDS_FILE(json_path, dataset_root, verbose=False) + ext_paths: dict[str, Path] = {ext: Path(p) for ext, p in bf_current.file.items()} + # Add same-stem companions that BIDS_FILE didn't pick up. json_path.name + # ends with e.g. "…_dwi.json" — strip the trailing ".json" to get the + # stem the .bval / .bvec / other companions share. + stem = json_path.name.removesuffix(".json") + for sibling in json_path.parent.iterdir(): + if not sibling.is_file(): + continue + if sibling.name == json_path.name: + continue + if not sibling.name.startswith(stem + "."): + continue + ext = sibling.name[len(stem) + 1 :] + ext_paths.setdefault(ext, sibling) + # Collision handling — resolve on the JSON pair (JSON header is the smallest + # authoritative sidecar, ideal for hashing). Then apply the same stem + # decision to every extension in the family. + action = "move" + final_json_path = new_json_path + if not dry_run: + source_json = ext_paths.get("json", json_path) + final_json_path, action = _resolve_stem_collision(new_json_path, Path(source_json)) + if action == "duplicate": + logger.on_neutral(f"duplicate content; removing redundant source {source_json}") + for src in ext_paths.values(): + Path(src).unlink(missing_ok=True) + return [] + if action == "run": + logger.on_warning(f"target {new_json_path} exists with different content; using {final_json_path}") + moves: list[tuple[Path, Path]] = [] + new_stem = str(final_json_path).removesuffix(".json") + for ext, src in ext_paths.items(): + dst = Path(f"{new_stem}.{ext}") + if not Path(src).exists(): + continue + moves.append((Path(src), dst)) + if dry_run: + return moves + for src, dst in moves: + dst.parent.mkdir(parents=True, exist_ok=True) + if dst.exists() and dst.resolve() != src.resolve(): + # Second-layer safety: after the stem-level resolve above one of + # the per-extension companions may still hit an unrelated + # existing file. Fall back to the old skip-with-warning here so + # data is never silently overwritten. + logger.on_warning(f"target {dst} already exists; skipping {src}") + continue + src.rename(dst) + return moves + + +_SEQU_ENTITY_RE = re.compile(r"_sequ-([^_\s.]+)") +_SES_ENTITY_RE = re.compile(r"_ses-([^_\s.]+)") +_RUN_ENTITY_RE = re.compile(r"_run-([^_\s.]+)") + + +def _files_identical(a: Path, b: Path) -> bool: + """Cheap byte-identity check: same size then same SHA-256. + + Used before applying a ``_run-`` collision slot — when the two files + are literally the same content there's nothing to keep, and we drop the + source instead of proliferating ``_run-2`` copies of identical bytes. + """ + try: + if a.stat().st_size != b.stat().st_size: + return False + except OSError: + return False + import hashlib + + h_a, h_b = hashlib.sha256(), hashlib.sha256() + for path, h in ((a, h_a), (b, h_b)): + try: + with path.open("rb") as fh: + for chunk in iter(lambda: fh.read(1 << 20), b""): + h.update(chunk) + except OSError: + return False + return h_a.digest() == h_b.digest() + + +def _insert_run_entity(bids_path: Path, n: int) -> Path: + """Return *bids_path* with a ``_run-`` entity inserted after ``_sequ-``. + + ``run`` is a legal BIDS entity value used to disambiguate repeated + acquisitions. We keep the rest of the filename identical and only splice + in ``_run-`` right after the ``_sequ-`` group, so downstream BIDS + parsers still see the same sequence number and pair it with a + disambiguating run index. + """ + name = bids_path.name + m = _SEQU_ENTITY_RE.search(name) + if not m: + # Fallback: prepend right before the trailing "_.". + stem, sep, tail = name.rpartition("_") + return bids_path.with_name(f"{stem}_run-{n}{sep}{tail}") if sep else bids_path + end = m.end() + return bids_path.with_name(name[:end] + f"_run-{n}" + name[end:]) + + +def _resolve_stem_collision(new_json_path: Path, source_json: Path) -> tuple[Path, str]: + """Pick a non-colliding target JSON path, or signal "drop as duplicate". + + Returns ``(final_json_path, action)`` where ``action`` is: + * ``"move"`` — target free, use ``new_json_path`` as-is; + * ``"run"`` — target taken with *different* content, use the + returned run-N variant instead; + * ``"duplicate"`` — target taken with *identical* content; caller + should drop the source rather than keep two copies. In this case + ``final_json_path`` is still ``new_json_path`` for the log. + """ + if not new_json_path.exists() or new_json_path.resolve() == source_json.resolve(): + return new_json_path, "move" + if _files_identical(new_json_path, source_json): + return new_json_path, "duplicate" + # Different content — walk `_run-2`, `_run-3`, … until we find a free slot. + n = 2 + while True: + candidate = _insert_run_entity(new_json_path, n) + if not candidate.exists(): + return candidate, "run" + n += 1 + if n > 999: + return new_json_path, "move" # give up; the outer skip guard kicks in + + +def _leftover_move_plan( + scan_root: Path, + subject_map: dict[str, str], + parent: str | None, +) -> list[tuple[Path, Path]]: + """Compute (src, dst) moves for files left behind after the JSON pass. + + After the JSON-driven Pass 1, an old ``sub-/ses-/`` folder can + still hold companion files whose primary sidecar has already migrated — + typical culprits are ``.bval`` / ``.bvec`` sitting apart from their + ``_dwi.nii.gz`` and DWI-derived ``_dwi_ADC.nii.gz`` maps. Move them under + ``sub-/`` too, preferring to co-locate with their new-side twin + (matched by ``sequ-``) so BIDS stem-linkage stays intact. Files that + can't be twinned drop at the session level of the new subject folder as + a safe fallback. + """ + del parent # kept for signature symmetry; scan_root already resolves it + moves: list[tuple[Path, Path]] = [] + for old_sub, new_sub in subject_map.items(): + if old_sub == new_sub: + continue + old_dir = scan_root / f"sub-{old_sub}" + new_dir = scan_root / f"sub-{new_sub}" + if not old_dir.is_dir(): + continue + for src in old_dir.rglob("*"): + if not src.is_file(): + continue + # Skip anything a re-run of the JSON pass would handle (we don't + # want to race the pass 1 output here). + if src.suffix == ".json": + continue + m = _SEQU_ENTITY_RE.search(src.name) + sequ = m.group(1) if m else None + m_ses = _SES_ENTITY_RE.search(src.name) + ses = m_ses.group(1) if m_ses else None + twin_stem: str | None = None + target_dir: Path | None = None + if sequ and new_dir.is_dir(): + for twin in new_dir.rglob(f"*sequ-{sequ}*.json"): + if not twin.is_file(): + continue + # Require (session, sequ) both to match. `sequ-` on its + # own repeats across sessions in longitudinal studies and + # would otherwise glue a 2022 DWI onto a 2025 SWI just + # because both happen to be scanner-slot 601. + twin_ses_m = _SES_ENTITY_RE.search(twin.name) + twin_ses = twin_ses_m.group(1) if twin_ses_m else None + if ses is not None and twin_ses is not None and ses != twin_ses: + continue + twin_stem = twin.name.removesuffix(".json") + target_dir = twin.parent + break + if twin_stem is not None and target_dir is not None: + # Extract the orphan's tail after the sequ- segment. That + # tail carries the old format label plus any trailing suffix + # (`_ADC`, extensions like `.bval` / `.bvec` / `.nii.gz`). + # Replacing the twin's stem preserves the format-label change + # while keeping DWI derivatives glued to the twin. + seq_marker = f"_sequ-{sequ}" + idx = src.name.find(seq_marker) + if idx == -1: + continue + after_sequ = src.name[idx + len(seq_marker) :] + # after_sequ starts with either `_...` or `.`. + # Strip the leading `_` so `_dwi.bval` becomes + # `.bval` and `_dwi_ADC.nii.gz` becomes `_ADC.nii.gz` — both + # then splice cleanly onto the twin stem. + if after_sequ.startswith("_"): + body, sep, rest = after_sequ.partition(".") + old_fmt_parts = body.split("_", 2) + # body = "_" or "__" + if len(old_fmt_parts) >= 2: + suffix = ("_" + old_fmt_parts[2]) if len(old_fmt_parts) == 3 else "" + after_sequ = f"{suffix}.{rest}" if sep else suffix + new_name = twin_stem + after_sequ + dst = target_dir / new_name + else: + # Fallback: mirror the old ses-* folder under sub-/ and + # just swap the subject prefix in the filename. + try: + rel = src.relative_to(old_dir) + except ValueError: + continue + new_name = src.name.replace(f"sub-{old_sub}", f"sub-{new_sub}", 1) + dst = new_dir / rel.parent / new_name + if src == dst or dst.exists(): + continue + moves.append((src, dst)) + return moves + + +def _prune_empty_folders(root: Path) -> int: + """Remove now-empty ``sub-*`` / ``ses-*`` subtrees left by the rename.""" + removed = 0 + for folder in sorted((p for p in root.rglob("*") if p.is_dir()), reverse=True): + rel = folder.relative_to(root) + if not rel.parts: + continue + try: + folder.rmdir() + removed += 1 + except OSError: + pass + return removed + + +def rerun_bids_naming( + dataset_root: Path | str, + parent: str | None = "rawdata", + subject_prefix: str | None = None, + subject_number_width: int = 3, + dry_run: bool = True, + info_dir: str = "info", + session: bool = True, + make_subject_chunks: int = 0, + verbose: bool = True, +) -> dict[str, list[tuple[Path, Path]]]: + """Rename an already-extracted BIDS dataset from its sidecar JSONs. + + Walks ``//`` (or ```` if ``parent`` is + ``None`` / empty — matches datasets where ``sub-*`` folders sit directly + under the dataset root), builds an ``old_sub -> new_sub`` map, then for + every JSON sidecar re-derives its target BIDS path via + :func:`extract_keys_from_json` + :func:`_generate_bids_path` and renames + each file family to the new path. + + Args: + dataset_root: Dataset root that contains the ``parent`` folder (or + the ``sub-*`` folders themselves). + parent: BIDS-style parent folder name (typically ``"rawdata"``). Pass + ``None`` or ``""`` for flat datasets like ``TOF_MPRAGE/sub-*``. + subject_prefix: If given, reassign every subject to + ``{prefix}{n:0Xd}`` starting from 1 (see ``subject_number_width``). + Otherwise sanitise the existing subject ids in place (strip ``_`` + and leading ``-``). + subject_number_width: Zero-padding width for the numeric suffix in + ``subject_prefix`` mode. Default 3 → ``ID001`` … ``ID999``. + dry_run: When ``True`` (default) plan the moves and log them but do + not touch disk. Flip to ``False`` to actually rename. + info_dir: Directory (relative to ``dataset_root``) where the + translation table lands. Defaults to ``"info"`` → written as + ``/info/subject_map.tsv``. + session: Forwarded to :func:`extract_keys_from_json` — populate the + ``ses`` entity from ``StudyDate`` when the sidecar lacks one. + make_subject_chunks: Forwarded to :func:`_generate_bids_path` (adds a + sub-folder built from the first N chars of the subject id). + verbose: Log every planned move. + + Returns: + Dict with keys ``"moves"`` (list of ``(src, dst)`` tuples), and + ``"mapping_file"`` (path of the written translation table). + """ + dataset_root = Path(dataset_root) + parent_norm = parent or "" + scan_root = dataset_root / parent_norm if parent_norm else dataset_root + if not scan_root.exists(): + raise FileNotFoundError(scan_root) + + subject_folders = _list_subject_folders(scan_root) + if not subject_folders: + logger.on_warning(f"No sub-* folders found under {scan_root}; nothing to do.") + return {"moves": [], "mapping_file": None} + + existing_map = _read_existing_subject_map(dataset_root, info_dir) + subject_map = _build_subject_map(subject_folders, subject_prefix, subject_number_width, existing_map) + logger.on_neutral( + f"Renaming {len(subject_folders)} subject(s) " + f"({'numeric ' + (subject_prefix or '') if subject_prefix else 'sanitising in place'}); " + f"dry_run={dry_run}." + ) + + all_moves: list[tuple[Path, Path]] = [] + for json_path in _iter_json_sidecars(scan_root): + # Recover the old sub id from the parent folder name — the JSON's + # own `PatientID` field is not authoritative once we start remapping. + try: + sub_folder = next(p for p in json_path.parents if p.name.startswith("sub-")) + except StopIteration: + continue + old_sub = _current_sub_id(sub_folder) + new_sub = subject_map.get(old_sub, _sanitize_entity_value(old_sub)) + try: + new_json_path, _keys = _new_bids_path_for( + json_path, + dataset_root, + parent_norm, + new_sub, + session=session, + make_subject_chunks=make_subject_chunks, + ) + except Exception as e: # noqa: BLE001 + logger.on_warning(f"Cannot re-derive BIDS name for {json_path.name}: {type(e).__name__}: {e}") + continue + if new_json_path == json_path: + continue + moves = _rename_family(json_path, new_json_path, dataset_root, dry_run=dry_run) + if verbose: + for src, dst in moves: + logger.on_neutral(f"{'[dry]' if dry_run else '[mv ]'} {src.relative_to(dataset_root)} -> {dst.relative_to(dataset_root)}") + all_moves.extend(moves) + + # Pass 2 — sweep orphan companions (.bval / .bvec / _ADC.nii.gz / other + # non-sidecar files) that Pass 1 didn't see because their JSON already + # moved on a previous run or they never had one. + leftover = _leftover_move_plan(scan_root, subject_map, parent_norm) + for src, dst in leftover: + if verbose: + logger.on_neutral(f"{'[dry]' if dry_run else '[lft]'} {src.relative_to(dataset_root)} -> {dst.relative_to(dataset_root)}") + if not dry_run: + dst.parent.mkdir(parents=True, exist_ok=True) + if dst.exists() and dst.resolve() != src.resolve(): + logger.on_warning(f"target {dst} already exists; skipping {src}") + continue + src.rename(dst) + all_moves.append((src, dst)) + + mapping_file = _write_subject_map(dataset_root, info_dir, subject_map) if not dry_run else None + if not dry_run: + removed = _prune_empty_folders(scan_root) + if removed: + logger.on_neutral(f"Pruned {removed} empty folder(s) after rename.") + logger.on_neutral(f"Planned {len(all_moves)} file move(s); mapping table: {mapping_file}") + return {"moves": all_moves, "mapping_file": mapping_file} + + +if __name__ == "__main__": + import argparse + + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("dataset_root", type=Path) + ap.add_argument("--parent", default="rawdata", help="'' / --parent '' for flat datasets") + ap.add_argument("--subject-prefix", default=None, help='e.g. "ID" for sub-ID001, sub-ID002, ...') + ap.add_argument("--subject-number-width", type=int, default=3) + ap.add_argument("--info-dir", default="info") + ap.add_argument("--make-subject-chunks", type=int, default=0) + ap.add_argument("--no-session", dest="session", action="store_false") + ap.add_argument("--apply", action="store_true", help="Actually perform the moves (default is dry-run).") + args = ap.parse_args() + + rerun_bids_naming( + args.dataset_root, + parent=args.parent or None, + subject_prefix=args.subject_prefix, + subject_number_width=args.subject_number_width, + dry_run=not args.apply, + info_dir=args.info_dir, + session=args.session, + make_subject_chunks=args.make_subject_chunks, + ) diff --git a/TPTBox/core/dicom/xa_pairing.py b/TPTBox/core/dicom/xa_pairing.py new file mode 100644 index 0000000..77cf383 --- /dev/null +++ b/TPTBox/core/dicom/xa_pairing.py @@ -0,0 +1,278 @@ +"""Post-extract pairing of biplane X-ray Angiography series. + +Biplane angio runs are exported as two separate DICOM series (A and B plane) +that share ``StudyInstanceUID`` and ``AcquisitionTime`` but carry different +``SeriesNumber`` values — one per plane. After :func:`extract_dicom_folder` +runs, that difference propagates into distinct ``sequ-`` entities on the +two BIDS filenames, which makes the two planes look like independent runs. + +This module rewires the pair so both planes carry the A-side's +``sequ-`` value — they now differ only by ``acq-A`` vs ``acq-B`` and +downstream tools that group by ``(sub, ses, sequ)`` see the biplane run as +one physical acquisition. + +Optionally also flips the Z-axis affine of every XA NIfTI to work around a +Siemens quirk that ships the volume upside-down (byte-identical voxel data, +just a sign flip on the third affine column). Off by default. +""" + +from __future__ import annotations + +import argparse +import json +from collections import defaultdict +from pathlib import Path + +import nibabel + +from TPTBox import BIDS_FILE, BIDS_Global_info, Print_Logger + +logger = Print_Logger() + +# Sidecar-JSON key we set after flipping the affine Z-column so subsequent +# runs can detect the file has already been corrected and skip it. Keeping +# state on the sidecar (rather than a separate ledger) survives file moves +# by the renamer and is what a re-extract from DICOMs would overwrite. +_Z_FLIP_MARKER = "TPTBoxXAZFlipped" + +# BIDS formats that carry XA-family payload. `.filter_format` accepts these +# values verbatim as the series's `bids_format`. +_XA_FORMATS: tuple[str, ...] = ("XA", "DSA", "DSA3D", "3DRA", "fluroscopy", "subtraction") + + +def _norm_time(t) -> str: + """Normalise a DICOM TM value (``HHMMSS.ffffff`` or ``HH:MM:SS.ffffff``).""" + if t is None: + return "" + s = str(t) + if ":" in s: + try: + hh, mm, rest = s.split(":", maxsplit=2) + except ValueError: + return s + else: + hh, mm, rest = s[:2], s[2:4], s[4:] or "0" + try: + return f"{int(hh):02d}:{int(mm):02d}:{float(rest):09.6f}" + except (ValueError, TypeError): + return s + + +def _bucket_key(json_obj: dict) -> tuple: + """Group a series by (StudyInstanceUID, SeriesNumber-family, AcquisitionTime). + + Both planes of a biplane run share ``StudyInstanceUID`` and + ``AcquisitionTime`` down to milliseconds; ``SeriesNumber`` differs by one + (A first, B second). Bucketing on the study UID + acquisition time is + tight enough to catch the pair without false positives. + """ + return ( + json_obj.get("StudyInstanceUID"), + _norm_time(json_obj.get("AcquisitionTime")), + ) + + +def _plane(bf: BIDS_FILE) -> str | None: + """Return the ``acq`` entity's plane label if it is ``A`` or ``B``.""" + val = bf.get("acq") + if val in ("A", "B"): + return val + return None + + +def _sidecar_for(nii_path: Path) -> Path: + """Return the ``.json`` sidecar path that pairs with a ``.nii.gz`` file.""" + name = nii_path.name + if name.endswith(".nii.gz"): + return nii_path.with_name(name[: -len(".nii.gz")] + ".json") + return nii_path.with_suffix(".json") + + +def _flip_marker_set(sidecar: Path) -> bool: + """True when the sidecar already records that its NIfTI was Z-flipped.""" + if not sidecar.is_file(): + return False + try: + j = json.loads(sidecar.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError, UnicodeDecodeError): + return False + return bool(j.get(_Z_FLIP_MARKER)) + + +def _write_flip_marker(sidecar: Path) -> None: + """Set ``_Z_FLIP_MARKER=true`` on the sidecar; no-op if it doesn't exist.""" + if not sidecar.is_file(): + return + try: + j = json.loads(sidecar.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError, UnicodeDecodeError): + return + j[_Z_FLIP_MARKER] = True + try: + sidecar.write_text(json.dumps(j, indent=4), encoding="utf-8") + except OSError as e: + logger.on_warning(f"Could not persist flip marker in {sidecar}: {e}") + + +def _flip_z_affine(path: Path) -> None: + """Negate the third affine column of the NIfTI at *path* in place. + + Pixel data is untouched — only the affine's Z-column (and its qform / + sform copies) get sign-flipped, which fixes the upside-down display on + Siemens XA sources. Idempotency is handled by the caller via + :func:`_flip_marker_set` / :func:`_write_flip_marker` on the sidecar — + a double-run of the pairing command with ``--flip-z`` would otherwise + negate twice and land back on the original orientation. + """ + img = nibabel.load(str(path)) + aff = img.affine.copy() + aff[:3, 2] = -aff[:3, 2] + out = nibabel.Nifti1Image(img.dataobj, aff, header=img.header) + out.set_qform(aff) + out.set_sform(aff) + nibabel.save(out, str(path)) + + +def pair_biplane_xa( + dataset_root: Path | str, + parent: str = "rawdata", + dry_run: bool = True, + flip_z: bool = False, + verbose: bool = True, +) -> dict: + """Rewrite biplane B-plane files to share A-plane's ``sequ`` value. + + Walks every XA-family series under ``//``, buckets + by ``(StudyInstanceUID, AcquisitionTime)``, and for each bucket that + contains both an ``acq-A`` and an ``acq-B`` member: if the two carry + different ``sequ-`` values, rename the B-side to use the A-side's + ``sequ`` (via :meth:`BIDS_FILE.rename_files`). All extensions in the + family follow the rename. + + When ``flip_z`` is set, every XA NIfTI seen (both A and B) also gets + its affine's third column negated — no-op on non-Siemens data whose + display was already right-side-up. + + Args: + dataset_root: Dataset root that contains ```` — same value + ``BIDS_Global_info`` would take. + parent: BIDS parent folder (``rawdata`` etc.). + dry_run: Default True — plan the moves and print them, don't touch + disk. Set False to apply. + flip_z: If True, negate every XA NIfTI's affine Z-column. Applied + to the file at its FINAL location (post-rename). + verbose: Log every planned rename. + + Returns: + Dict with ``"pairs_matched"`` (count of (A,B) buckets found), + ``"renames"`` (list of ``(src, dst)`` tuples), and ``"flipped"`` + (list of paths whose affine was flipped). + """ + dataset_root = Path(dataset_root) + bgi = BIDS_Global_info(datasets=[dataset_root], parents=[parent]) + buckets_by_subject: dict[str, dict[tuple, dict[str, list[BIDS_FILE]]]] = {} + for sub_name, subj in bgi.enumerate_subjects(): + q = subj.new_query(flatten=True) + q.filter_format(list(_XA_FORMATS)) + subj_buckets: dict[tuple, dict[str, list[BIDS_FILE]]] = defaultdict(lambda: {"A": [], "B": []}) + for bf in q.loop_list(): + plane = _plane(bf) + if plane is None: + continue + try: + j = bf.open_json() + except Exception: # noqa: BLE001 + continue + key = _bucket_key(j) + subj_buckets[key][plane].append(bf) + buckets_by_subject[sub_name] = subj_buckets + + renames: list[tuple[Path, Path]] = [] + flipped: list[Path] = [] + pairs_matched = 0 + for sub_name, subj_buckets in buckets_by_subject.items(): + for planes in subj_buckets.values(): + a_side = planes["A"] + b_side = planes["B"] + if not a_side or not b_side: + continue + pairs_matched += 1 + # A-side's `sequ` is authoritative — smaller SeriesNumber ships + # first on Siemens/Philips biplane exports. Pick lowest if + # multiple A-members exist (rare — repeated runs at same time). + a_sequ = min({str(bf.get("sequ") or "") for bf in a_side}) + if not a_sequ: + continue + for bf in b_side: + b_sequ = str(bf.get("sequ") or "") + if b_sequ == a_sequ: + continue + # Build the target path with A's sequ; keep all other entities. + new_bids = bf.get_changed_bids(info={"sequ": a_sequ}, non_strict_mode=True) + src_paths = {ext: Path(p) for ext, p in bf.file.items()} + new_nii = Path(new_bids.file.get("nii.gz", "")) + new_stem = str(new_nii).removesuffix(".nii.gz") + for ext, src in src_paths.items(): + dst = Path(f"{new_stem}.{ext}") + if src == dst or not src.exists(): + continue + if verbose: + logger.on_neutral( + f"{'[dry]' if dry_run else '[mv ]'} sub={sub_name} B→A pairing " + f"{src.relative_to(dataset_root)} -> {dst.relative_to(dataset_root)}" + ) + renames.append((src, dst)) + if not dry_run: + dst.parent.mkdir(parents=True, exist_ok=True) + if dst.exists() and dst.resolve() != src.resolve(): + logger.on_warning(f"target {dst} already exists; skipping {src}") + continue + src.rename(dst) + + if flip_z: + # Re-scan after renames so we hit files at their final paths. + bgi2 = BIDS_Global_info(datasets=[dataset_root], parents=[parent]) + skipped_already_flipped = 0 + for _sub, subj in bgi2.enumerate_subjects(): + q = subj.new_query(flatten=True) + q.filter_format(list(_XA_FORMATS)) + for bf in q.loop_list(): + if _plane(bf) is None: + continue + nii = bf.file.get("nii.gz") + if nii is None or not Path(nii).exists(): + continue + sidecar = _sidecar_for(Path(nii)) + # Idempotency: a sidecar that already carries our marker was + # flipped on a previous run; flipping again would undo it. + if _flip_marker_set(sidecar): + skipped_already_flipped += 1 + continue + if verbose: + logger.on_neutral(f"{'[dry]' if dry_run else '[fz ]'} flip-Z {Path(nii).relative_to(dataset_root)}") + if not dry_run: + _flip_z_affine(Path(nii)) + _write_flip_marker(sidecar) + flipped.append(Path(nii)) + if skipped_already_flipped: + logger.on_neutral(f"flip-Z: skipped {skipped_already_flipped} file(s) that already carry the marker.") + + logger.on_neutral( + f"Biplane pairs matched: {pairs_matched}; renames planned: {len(renames)}; flipped: {len(flipped)} (dry_run={dry_run})" + ) + return {"pairs_matched": pairs_matched, "renames": renames, "flipped": flipped} + + +if __name__ == "__main__": + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("dataset_root", type=Path) + ap.add_argument("--parent", default="rawdata") + ap.add_argument("--flip-z", action="store_true", help="Also negate the Z-column of every XA NIfTI's affine (Siemens quirk).") + ap.add_argument("--apply", action="store_true", help="Perform the moves; default is dry-run.") + args = ap.parse_args() + pair_biplane_xa( + args.dataset_root, + parent=args.parent, + dry_run=not args.apply, + flip_z=args.flip_z, + ) diff --git a/TPTBox/core/internal/train_nnUnet/_prep_ds.py b/TPTBox/core/internal/train_nnUnet/_prep_ds.py index 2328da6..451e3af 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 c4dddf0..bf9a4da 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 a153924..a9677dc 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,99 @@ 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 +249,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 +275,23 @@ 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 +414,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 +481,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() diff --git a/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py b/TPTBox/segmentation/nnUnet_utils/sliding_window_prediction.py index 88708e5..400057a 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/README.md b/TPTBox/spine/spinestats/README.md index 7e04fed..5afa6ae 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` @@ -333,15 +468,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: diff --git a/TPTBox/spine/spinestats/_load_nako.py b/TPTBox/spine/spinestats/_load_nako.py index 9b77f31..b0462ee 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", @@ -234,8 +235,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/_load_nako_wh.py b/TPTBox/spine/spinestats/_load_nako_wh.py new file mode 100644 index 0000000..3925965 --- /dev/null +++ b/TPTBox/spine/spinestats/_load_nako_wh.py @@ -0,0 +1,1422 @@ +import json +import os +import tempfile +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") + +_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. + + 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. + """ + if _NON_INTERACTIVE: + return "__skip__", None + 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") + 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 + 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: + """Verify that every expected base image is either present or confirmed missing. + + 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 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.") + 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 + + +_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. + + 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 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[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 + + 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"[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): + """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 [] + 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: + 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="/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. + + 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 + + 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", # 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: + 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 + + # 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") + # 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: + # 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", 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 # 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) + 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) + # 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", + "vibe_part-fat", + "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: + 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_by_key: dict = {} + 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 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 + 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 + if corrected_index: + _apply_corrections_to_subj_dict(str(sub), subj_dict, corrected_index) + verify_missing_images(cache, sub, subj_dict) + yield subj_dict + + +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/", +): + 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 + # 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) + ] + 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 or bf == "": + 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", + ) + # 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: + """Dry-run a single subject's hard-link plan. + + 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 t2w in (d.get("t2w_chunk") or {}).values(): + 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]: + """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) + except Exception as e: # noqa: BLE001 + return nii_path, f"error: {type(e).__name__}: {e}" + else: + return nii_path, "ok" + + +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 + # 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 + try: + 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(): + 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 _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): + 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, 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).") + args = parser.parse_args() + 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: + corrected = load_corrected_index() + for d in loop_over_repaired_nako(test=test, corrected_index=corrected, skip_subject=is_hard_linked): + hard_link(d) diff --git a/TPTBox/spine/spinestats/_qc_report.py b/TPTBox/spine/spinestats/_qc_report.py new file mode 100644 index 0000000..54d3b12 --- /dev/null +++ b/TPTBox/spine/spinestats/_qc_report.py @@ -0,0 +1,231 @@ +"""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 4ba2934..236d189 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, } @@ -145,13 +157,15 @@ 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, 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,12 +316,12 @@ 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 - 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 @@ -281,24 +342,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) @@ -306,7 +378,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) @@ -317,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) @@ -345,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``. @@ -399,30 +585,42 @@ 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")} + # 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]] = [] - 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( @@ -430,6 +628,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 @@ -438,21 +637,65 @@ 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: - 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) + 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: @@ -461,12 +704,13 @@ 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() + _log(f"heartbeat: seen={len(seen)} vertebra_rows={len(vertebra_rows)} ivd_rows={len(ivd_rows)}") class ExcelCollector: @@ -488,11 +732,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 @@ -502,7 +748,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() @@ -512,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) @@ -572,6 +833,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 @@ -579,14 +841,15 @@ 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 = 40 # set >1 to parallelize + N_CPUS = 1 # set >1 to parallelize OVERRIDE = False aggregate = True do_not_update = False test = False + collector: ExcelCollector | None = None if aggregate: collector = ExcelCollector(out_folder=OUT_FOLDER) collector.start() @@ -595,18 +858,20 @@ 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) 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) @@ -617,8 +882,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() @@ -637,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 acff529..2cc16a0 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,19 +250,119 @@ 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 -`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: diff --git a/TPTBox/spine/spinestats/angles.py b/TPTBox/spine/spinestats/angles.py index 3a6af2d..606f56a 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/curvature.py b/TPTBox/spine/spinestats/curvature.py new file mode 100644 index 0000000..8f8f267 --- /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/measure_ivd_and_vertebra_geometry.py b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py index f16ce48..1d08bae 100644 --- a/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py +++ b/TPTBox/spine/spinestats/measure_ivd_and_vertebra_geometry.py @@ -177,20 +177,27 @@ 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 +447,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 +520,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 +545,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 +635,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 +659,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 +686,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 +725,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 +775,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 +820,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 +846,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 diff --git a/TPTBox/spine/spinestats/pelvic_parameters.py b/TPTBox/spine/spinestats/pelvic_parameters.py new file mode 100644 index 0000000..9f449ec --- /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 diff --git a/TPTBox/spine/spinestats/torso_vat_sat.py b/TPTBox/spine/spinestats/torso_vat_sat.py index c165011..ad6fa4f 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}"]