From a34fc1d26d1da32ba6b42d846c7bcfa2ded296f8 Mon Sep 17 00:00:00 2001 From: Stephen Aylward Date: Sat, 8 Aug 2026 11:33:44 -0400 Subject: [PATCH 1/3] ENH: Per-organ tutorial parameters and distance-map registration fixes Introduce tutorials/parameters_lung_ct_dirlab.py and tutorials/parameters_heart_ct_kcl.py as the single source for each use case's mask dilation, distance-map saturation radius, PCA component counts, Greedy iteration schedule, segmenter class, and (heart only) the interior chamber label ids. Every tutorial that rasterizes or registers a distance map now reads the same values, so the maps a network is finetuned on match the maps it later infers over. No paths live in these modules; each tutorial keeps its own inputs and outputs. Add tutorials/tutorial_02_heart_distancemap_finetune_icon.py, which finetunes uniGradICON on heart distance maps built from the Duke-Heart-4DLabelmaps labelmaps with the chambers excluded. The heart needs its own run rather than reusing the lung weights: its registration mask is much tighter, so its distance maps saturate over a shorter radius and do not share an intensity distribution with the lung ones. Library fixes: - transform_tools.transform_image gains an explicit background_value. Resampling previously fell back to ITK's default of 0, which for CT is water, not air, so pre-warped moving images carried a false soft-tissue shell wherever they had no data. register_images_base now fills with -1000 HU for CT (exactly uniGradICON's window floor) on the image warp only; masks and labelmaps keep 0. - register_models_distance_maps composed the Greedy and ICON transforms in the wrong order. ITK's CompositeTransform applies back to front, so the residual must be added last in the forward and first in the inverse. Also drops two unconditional debug_*.nii.gz writes that crashed when mask_dilation_mm was 0. - register_from restores moving_image and clears moving_image_registered, so a registrar can be reused after an initialized run. - workflow_fit_statistical_model_to_patient grids the PCA field on a template-frame reference image rather than the patient image, and pads physically from spacing. - Default registrar for intensity registration switches from RegisterImagesGreedyICON/ICON to RegisterImagesGreedy in register_time_series_images and workflow_convert_image_to_usd. Distance-map registration keeps ICON, now with the finetuned weights. - convert_vtk_to_usd validates object names as USD identifiers and rejects duplicates; workflow_convert_vtk_to_usd wraps raw vtkDataSet. Rename number_of_components / number_of_modes to number_of_pca_components throughout the workflows and tutorials. tutorial_02_lung_finetune_icon now writes difference images (fixed minus registered) instead of the resampled volumes, and reports the chain's Greedy-stage-only score as its own row. The chain remains unconditional: on DIR-Lab, ICON's 175^3 residual grid is about 1.4 mm over the FOV, coarser than the 1.10 mm Greedy already achieves, so it cannot refine and the tutorial reports that honestly. tutorial_02_lung_distancemap_finetune_icon restricts its cached labelmaps to the lung labels. They previously held all 97 whole-body classes, and uniGradICON's Dice loss one-hots every shared class at 175^3 by batch 4, which saturated GPU memory. Baselines for the slow and GPU buckets will need refreshing: the composition-order fix and the PCA field grid change alter registration output. --- data/Duke-Heart-4DLabelmaps/.gitignore | 1 + data/Duke-Heart-4DLabelmaps/README.md | 31 ++ docs/api/registration/chained.rst | 4 +- docs/tutorials.rst | 50 +- pyproject.toml | 3 + src/physiotwin4d/register_images_base.py | 52 +- src/physiotwin4d/register_images_chain.py | 11 + .../register_models_distance_maps.py | 4 +- .../register_time_series_images.py | 8 +- src/physiotwin4d/transform_tools.py | 18 +- .../workflow_convert_image_to_usd.py | 9 +- ...rkflow_fit_statistical_model_to_patient.py | 24 +- .../workflow_reconstruct_highres_4d_ct.py | 3 +- tests/test_workflow_convert_image_to_usd.py | 9 +- ...test_workflow_reconstruct_highres_4d_ct.py | 9 +- tutorials/README.md | 3 +- tutorials/parameters_heart_ct_kcl.py | 89 ++++ tutorials/parameters_lung_ct_dirlab.py | 78 +++ .../tutorial_01_heart_gated_ct_to_usd.py | 17 +- tutorials/tutorial_01_lung_gated_ct_to_usd.py | 17 +- ...02_duke_heart_distancemap_finetune_icon.py | 471 ++++++++++++++++++ ...orial_02_lung_distancemap_finetune_icon.py | 45 +- tutorials/tutorial_02_lung_finetune_icon.py | 226 ++++++++- ...rial_03_heart_reconstruct_highres_4d_ct.py | 7 +- ...orial_03_lung_reconstruct_highres_4d_ct.py | 37 +- ...orial_06_heart_create_statistical_model.py | 4 +- ...torial_06_lung_create_statistical_model.py | 5 +- ..._heart_fit_statistical_model_to_patient.py | 36 +- ...7_lung_fit_statistical_model_to_patient.py | 15 +- ...torial_08_lung_fit_model_to_4d_patients.py | 63 +-- 30 files changed, 1138 insertions(+), 211 deletions(-) create mode 100644 data/Duke-Heart-4DLabelmaps/.gitignore create mode 100644 data/Duke-Heart-4DLabelmaps/README.md create mode 100644 tutorials/parameters_heart_ct_kcl.py create mode 100644 tutorials/parameters_lung_ct_dirlab.py create mode 100644 tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py diff --git a/data/Duke-Heart-4DLabelmaps/.gitignore b/data/Duke-Heart-4DLabelmaps/.gitignore new file mode 100644 index 00000000..fb2d2365 --- /dev/null +++ b/data/Duke-Heart-4DLabelmaps/.gitignore @@ -0,0 +1 @@ +pm* diff --git a/data/Duke-Heart-4DLabelmaps/README.md b/data/Duke-Heart-4DLabelmaps/README.md new file mode 100644 index 00000000..1af9a9db --- /dev/null +++ b/data/Duke-Heart-4DLabelmaps/README.md @@ -0,0 +1,31 @@ +# Duke-Heart-4DLabelmaps + +Gated 4D cardiac labelmaps acquired at Duke University by Dr. Paul Segars. + +## Availability + +This dataset is **not currently available**. It is being considered for public +release; until that happens it cannot be downloaded, and it is not distributed +with this repository. + +## Effect on the tutorials + +Tutorials that depend on this dataset are named with a `duke_heart` prefix in +their organ field, for example: + +- `tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py` + +These `duke_heart` tutorials will not run without the data. Every other +tutorial uses a publicly available dataset and is unaffected — see +[../README.md](../README.md) for download instructions. + +Downstream tutorials that consume `duke_heart` outputs (such as the finetuned +distance-map ICON weights used by +`tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py`) fall back to +stock weights and still run, with reduced accuracy. + +## Expected layout + +When available, the data is expected under `data/Duke-Heart-4DLabelmaps/` as +one directory per case (`pm0002/`, `pm0003/`, ...), each holding one labelmap +per gated frame. diff --git a/docs/api/registration/chained.rst b/docs/api/registration/chained.rst index a07d89f7..e86038bc 100644 --- a/docs/api/registration/chained.rst +++ b/docs/api/registration/chained.rst @@ -9,8 +9,8 @@ Chained Image Registration Coarse-to-fine registration composes two registrars: a fast, robust method recovers the large motion, then a deformable method refines it. ``RegisterImagesChain`` is the general composition; -``RegisterImagesGreedyICON`` is the Greedy-then-ICON pairing the tutorials use -for 4D CT. +``RegisterImagesGreedyICON`` is the Greedy-then-ICON pairing, used by Tutorial 2 +and by the distance-map stage of the statistical-model fit. Both implement :class:`RegisterImagesBase`, so they drop into any workflow that takes a ``registration_method``. diff --git a/docs/tutorials.rst b/docs/tutorials.rst index 8c25dc3f..d505f601 100644 --- a/docs/tutorials.rst +++ b/docs/tutorials.rst @@ -11,7 +11,7 @@ Tutorials

PhysioTwin4D tutorials

From a CT scan to an animated digital twin

- Ten numbered stages across 16 runnable Python scripts. + Ten numbered stages across 17 runnable Python scripts. Each one drives the real workflow classes end-to-end on downloadable data, shows what it produced, and ends with the handful of constants to change so it runs on your own scans. @@ -152,7 +152,7 @@ Script Workflow :class:`~physiotwin4d.WorkflowConvertImageToUSD`, driving - :class:`~physiotwin4d.RegisterImagesICON` and a + :class:`~physiotwin4d.RegisterImagesGreedy` and a :class:`~physiotwin4d.SegmentAnatomyBase` subclass. Dataset @@ -161,9 +161,8 @@ Dataset registration reference. Requirements - GPU strongly recommended — ICON registers every phase against the - reference. Swap in :class:`~physiotwin4d.RegisterImagesGreedy` or - :class:`~physiotwin4d.RegisterImagesANTS` for CPU-only environments. + Greedy registers every phase against the reference on the CPU; a GPU is + still needed for segmentation. Preview .. figure:: assets/tutorial_01_heart_4d.gif @@ -219,11 +218,21 @@ Tutorial 2: Finetune ICON Registration Script ``tutorials/tutorial_02_lung_finetune_icon.py`` - ``tutorials/tutorial_02_lung_distancemap_finetune_icon.py`` — the + ``tutorials/tutorial_02_lung_distancemap_finetune_icon.py`` — the lung distance-map variant, which finetunes on distance maps rather than image intensities so the labelmap-to-labelmap stage of Tutorials 7 and 8 has in-distribution weights. + ``tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py`` — the same + for the heart, on Duke-Heart-4DLabelmaps. The heart needs its own run because + it registers with a much tighter mask than the lungs, so its distance maps + saturate over a shorter radius and do not share an intensity distribution + with lung ones. The per-organ values live in + ``tutorials/parameters_lung_ct_dirlab.py`` and + ``tutorials/parameters_heart_ct_kcl.py``. This is a ``duke_heart`` tutorial: + Duke-Heart-4DLabelmaps is not publicly available yet, so it cannot be run — + see ``data/Duke-Heart-4DLabelmaps/README.md``. + Workflow :class:`~physiotwin4d.WorkflowFinetuneICONRegistration`, then :class:`~physiotwin4d.RegisterImagesGreedy` and @@ -301,20 +310,15 @@ Script Workflow :class:`~physiotwin4d.WorkflowReconstructHighres4DCT` with - :class:`~physiotwin4d.RegisterImagesGreedyICON`. + :class:`~physiotwin4d.RegisterImagesGreedy`. Dataset Slicer-Heart-CT for the heart; DIR-Lab for the lung, which reconstructs against its T70 (end-exhale) phase — the same reference Tutorial 8 fits to. - The lung variant registers with **Tutorial 2's finetuned ICON weights** when - they exist, and logs a warning and falls back to the stock uniGradICON - weights when they do not. Requirements - GPU recommended. One coarse-to-fine registration per phase, greedy schedule - ``[30, 15, 7, 3]``. The lung variant enables mass preservation for - non-contrast CT; the heart variant does not and uses the stock uniGradICON - weights. + CPU is enough. One coarse-to-fine registration per phase, greedy schedule + ``[30, 15, 7, 3]``. Preview .. figure:: assets/Tutorial_03_heart_original.gif @@ -338,11 +342,8 @@ Preview Inner API usage .. code-block:: python - registration_method = RegisterImagesGreedyICON() - registration_method.greedy.set_number_of_iterations([30, 15, 7, 3]) - registration_method.icon.set_mass_preservation(True) - if icon_weights_path.exists(): # Tutorial 2 output, optional - registration_method.icon.set_weights_path(str(icon_weights_path)) + registration_method = RegisterImagesGreedy() + registration_method.set_number_of_iterations([30, 15, 7, 3]) workflow = WorkflowReconstructHighres4DCT( time_series_images=time_series, @@ -366,10 +367,7 @@ Outputs Adapt to your data Set ``case_glob`` and ``data_dir`` to your series and pick the reference - with ``reference_time_frame``. Point ``icon_weights_path`` at weights you - finetuned on your own cohort with Tutorial 2, or leave it missing to - register with the stock uniGradICON weights. If you have a separate - breath-hold or + with ``reference_time_frame``. If you have a separate breath-hold or contrast-enhanced volume, pass it as ``reference_image`` instead of one of the phases — that is what the workflow is really designed for. Tune ``number_of_iterations_greedy`` down for a fast smoke test. The saved @@ -657,9 +655,9 @@ Workflow the fitted surface through every other phase. Dataset - DIR-Lab, plus Tutorial 6 (lung)'s model. Tutorial 2's finetuned ICON weights - are used when present; without them the tutorial warns and registers with the - stock uniGradICON weights. + DIR-Lab, plus Tutorial 6 (lung)'s model. Tutorial 2's finetuned distance-map + ICON weights are used by the model fit when present; without them the + tutorial warns and fits with the stock uniGradICON weights. Requirements GPU required, and the heaviest registration workload in the set: one diff --git a/pyproject.toml b/pyproject.toml index ed8dba5c..aa20f911 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -325,8 +325,11 @@ module = [ # pre-commit mypy environment and ships no py.typed marker when installed. # mypy per-module patterns match whole components only, so each tutorial is # listed by name; add a line here when adding or renaming a tutorial. + "parameters_heart_ct_kcl", + "parameters_lung_ct_dirlab", "tutorial_01_heart_gated_ct_to_usd", "tutorial_01_lung_gated_ct_to_usd", + "tutorial_02_duke_heart_distancemap_finetune_icon", "tutorial_02_lung_distancemap_finetune_icon", "tutorial_02_lung_finetune_icon", "tutorial_03_heart_reconstruct_highres_4d_ct", diff --git a/src/physiotwin4d/register_images_base.py b/src/physiotwin4d/register_images_base.py index 7f3c1971..c148899c 100644 --- a/src/physiotwin4d/register_images_base.py +++ b/src/physiotwin4d/register_images_base.py @@ -19,6 +19,7 @@ from typing import Any, Optional, Union, cast import itk +import numpy as np from .labelmap_tools import LabelmapTools from .physiotwin4d_base import PhysioTwin4DBase @@ -108,6 +109,7 @@ def __init__(self, log_level: int | str = logging.INFO) -> None: self.moving_labelmap: Optional[itk.Image] = None self.mask_dilation_mm: float = 5.0 + self.prewarp_background_value: Optional[float] = None self.fast_mode: bool = False @@ -132,6 +134,39 @@ def set_modality(self, modality: str) -> None: """ self.modality = modality + def set_prewarp_background_value(self, background_value: float) -> None: + """Override the value a seeded registration's pre-warp writes off-grid. + + Args: + background_value: Intensity written where the fixed grid samples + outside the moving image. Leave unset to derive it from the + modality; see :meth:`_prewarp_background_value`. + """ + self.prewarp_background_value = background_value + + def _prewarp_background_value(self, moving_image: itk.Image) -> float: + """Return the intensity that means "no tissue" for the moving image. + + Pre-warping onto the fixed grid samples outside the moving image + wherever the two extents disagree. ITK's default fill of 0 is wrong for + an intensity image: in CT it is water, so the filled region reads as + soft tissue rather than air and any downstream similarity metric treats + it as structure to match. -1000 HU is also uniGradICON's ``ct_window`` + lower bound, so it normalizes to exactly the same value as true air. + + Args: + moving_image: Image being pre-warped. + + Returns: + The explicit override when set, -1000.0 for CT, otherwise the + moving image's own minimum intensity. + """ + if self.prewarp_background_value is not None: + return self.prewarp_background_value + if self.modality == "ct": + return -1000.0 + return float(np.min(itk.GetArrayViewFromImage(moving_image))) + def set_fixed_image(self, fixed_image: itk.Image) -> None: """Set the fixed/target image for registration. @@ -435,7 +470,10 @@ def _prewarp_moving( Returns: Tuple of the warped ``(image, mask, labelmap)``, the latter two None when not supplied. The mask and labelmap are warped with - nearest-neighbor interpolation to keep their discrete values. + nearest-neighbor interpolation to keep their discrete values, and + filled with 0 off-grid; the image is filled with + :meth:`_prewarp_background_value` instead, since 0 is a tissue + intensity rather than an absence of tissue. Raises: ValueError: If the fixed image has not been set. @@ -444,7 +482,11 @@ def _prewarp_moving( raise ValueError("Fixed image must be set before registration.") transform_tools = TransformTools() - self.log_info("Pre-warping moving data with the initial transform...") + background_value = self._prewarp_background_value(moving_image) + self.log_info( + "Pre-warping moving data with the initial transform (background %.1f)...", + background_value, + ) def _warp(image: Optional[itk.Image], nearest: bool) -> Optional[itk.Image]: if image is None: @@ -454,6 +496,7 @@ def _warp(image: Optional[itk.Image], nearest: bool) -> Optional[itk.Image]: initial_forward_transform, self.fixed_image, interpolation_method="nearest" if nearest else "linear", + background_value=0.0 if nearest else background_value, ) return ( @@ -479,7 +522,10 @@ def _compose_with_initial( Returns: The result dict with both transforms mapping between the *original* - moving image and the fixed image. + moving image and the fixed image. ``loss`` is passed through + unchanged, so it is the residual stage's loss measured against the + already pre-warped data -- not a loss for the composed transform, + and not comparable to the loss of a stage that started from scratch. """ transform_tools = TransformTools() diff --git a/src/physiotwin4d/register_images_chain.py b/src/physiotwin4d/register_images_chain.py index d05a719a..9a38f38b 100644 --- a/src/physiotwin4d/register_images_chain.py +++ b/src/physiotwin4d/register_images_chain.py @@ -28,6 +28,17 @@ class RegisterImagesChain(RegisterImagesBase): See :class:`RegisterImagesGreedyICON` for a named 2-stage convenience subclass (Greedy followed by ICON refinement). + Chaining is not free accuracy. Every stage's result is applied + unconditionally, so a refinement stage helps only when its own accuracy + floor is below the error the previous stage has already reached. A stage + whose deformation model is coarser than that error cannot resolve what is + left and acts as a low-pass perturbation, giving a slightly worse answer for + strictly more runtime. Compare each stage against the one before it on a + held-out metric rather than assuming the chain wins. + + ``result["loss"]`` is the *last* stage's loss, measured against data the + earlier stages already warped; it is not comparable to a single-stage loss. + Example: >>> chain = RegisterImagesChain([RegisterImagesGreedy(), RegisterImagesICON()]) >>> chain.set_fixed_image(fixed_image) diff --git a/src/physiotwin4d/register_models_distance_maps.py b/src/physiotwin4d/register_models_distance_maps.py index 43998b54..6f4e0f0c 100644 --- a/src/physiotwin4d/register_models_distance_maps.py +++ b/src/physiotwin4d/register_models_distance_maps.py @@ -135,7 +135,9 @@ def __init__( reference_image: ITK image providing coordinate frame (origin, spacing, direction) for mask generation. Typically the patient CT/MRI image. distance_squared_max: Maximum squared distance, in squared millimeters, - that the distance maps are normalized against. Default: 50.0 + that the distance maps are normalized against. It fixes their + intensity distribution, so it must match the value the ICON + weights in use were finetuned at. Default: 50.0 mask_dilation_mm: Dilation amount in millimeters for binary registration mask generation. Default: 20mm log_level: Logging level (default: logging.INFO) diff --git a/src/physiotwin4d/register_time_series_images.py b/src/physiotwin4d/register_time_series_images.py index 4da6f2e7..16a481ac 100644 --- a/src/physiotwin4d/register_time_series_images.py +++ b/src/physiotwin4d/register_time_series_images.py @@ -16,7 +16,7 @@ import itk from .register_images_base import RegisterImagesBase -from .register_images_greedy_icon import RegisterImagesGreedyICON +from .register_images_greedy import RegisterImagesGreedy from .transform_tools import TransformTools @@ -73,8 +73,7 @@ def __init__( Args: registration_method: Registration backend instance to use. - Defaults to a new RegisterImagesGreedyICON when None, with - its greedy stage configured to use an Affine transform. + Defaults to a new RegisterImagesGreedy when None. log_level: Logging level (default: logging.INFO) Raises: @@ -84,8 +83,7 @@ def __init__( super().__init__(log_level=log_level) if registration_method is None: - registration_method = RegisterImagesGreedyICON(log_level=log_level) - registration_method.greedy.set_transform_type("Affine") + registration_method = RegisterImagesGreedy(log_level=log_level) elif not isinstance(registration_method, RegisterImagesBase): raise TypeError( "registration_method must be a RegisterImagesBase instance or None" diff --git a/src/physiotwin4d/transform_tools.py b/src/physiotwin4d/transform_tools.py index b891156a..5f12bed2 100644 --- a/src/physiotwin4d/transform_tools.py +++ b/src/physiotwin4d/transform_tools.py @@ -12,7 +12,7 @@ """ import logging -from typing import Type, cast +from typing import Type, Union, cast import itk import numpy as np @@ -449,6 +449,7 @@ def transform_image( tfm: itk.Transform, reference_image: itk.image, interpolation_method: str = "linear", + background_value: float = 0.0, ) -> itk.image: """ Transform an ITK image using a specified transform and interpolation. @@ -467,6 +468,11 @@ def transform_image( - "linear": Linear interpolation (default, good for CT/MR) - "nearest": Nearest neighbor (preserves discrete values) - "sinc": Sinc interpolation (highest quality, slower) + background_value (float): Value written where the reference grid + samples outside the input image. Default 0.0, which is right for + labelmaps and masks; intensity images need the value that means + "no tissue" in their own units -- for CT that is -1000 HU (air), + not 0 HU (water). Returns: itk.image: The transformed image resampled to reference grid @@ -511,12 +517,22 @@ def transform_image( tfm, reference_image ) + # ITK's wrapping types DefaultPixelValue to the image's pixel type, and + # rejects a Python float for an integer image. + dtype = itk.GetArrayViewFromImage(img).dtype + default_pixel_value: Union[int, float] = ( + int(round(background_value)) + if np.issubdtype(dtype, np.integer) + else float(background_value) + ) + img_reg = itk.resample_image_filter( Input=img, Transform=dftfm, Interpolator=interpolator, ReferenceImage=reference_image, UseReferenceImage=True, + DefaultPixelValue=default_pixel_value, ) return img_reg diff --git a/src/physiotwin4d/workflow_convert_image_to_usd.py b/src/physiotwin4d/workflow_convert_image_to_usd.py index 1e2bae78..ca639889 100644 --- a/src/physiotwin4d/workflow_convert_image_to_usd.py +++ b/src/physiotwin4d/workflow_convert_image_to_usd.py @@ -20,7 +20,7 @@ from .image_tools import ImageTools from .physiotwin4d_base import PhysioTwin4DBase from .register_images_base import RegisterImagesBase -from .register_images_icon import RegisterImagesICON +from .register_images_greedy import RegisterImagesGreedy from .segment_anatomy_base import SegmentAnatomyBase from .segment_chest_total_segmentator_with_contrast import ( SegmentChestTotalSegmentatorWithContrast, @@ -41,7 +41,7 @@ class WorkflowConvertImageToUSD(PhysioTwin4DBase): instance. Configure backend-specific parameters (iteration counts, trim_branches, mass preservation, etc.) on the instance before passing it in. Defaults to :class:`SegmentChestTotalSegmentatorWithContrast` / - :class:`RegisterImagesICON` when omitted. + :class:`RegisterImagesGreedy` when omitted. """ def __init__( @@ -70,7 +70,7 @@ def __init__( backend instance. Defaults to a new :class:`SegmentChestTotalSegmentatorWithContrast` when None. registration_method (Optional[RegisterImagesBase]): Registration - backend instance. Defaults to a new :class:`RegisterImagesICON` + backend instance. Defaults to a new :class:`RegisterImagesGreedy` when None. A caller-supplied instance is mutated (fixed image/mask/modality) during :meth:`process` - pass a fresh instance per run unless intentionally reusing state. @@ -115,8 +115,7 @@ def __init__( self.segmenter: SegmentAnatomyBase = segmentation_method if registration_method is None: - registration_method = RegisterImagesICON(log_level=log_level) - registration_method.set_mass_preservation(False) + registration_method = RegisterImagesGreedy(log_level=log_level) elif not isinstance(registration_method, RegisterImagesBase): raise TypeError( "registration_method must be a RegisterImagesBase instance or None" diff --git a/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py b/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py index 30ec9e33..ea97d1ef 100644 --- a/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py +++ b/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py @@ -245,6 +245,7 @@ def __init__( # Parameters for labelmap and mask generation self.mask_dilation_mm: float = 10.0 # For binary registration mask generation + self.distancemap_squared_max: Optional[float] = None # Optional finetuned ICON checkpoint for the labelmap-to-labelmap stage self.l2l_icon_weights_path: Optional[str] = None @@ -300,6 +301,27 @@ def set_mask_dilation_mm(self, mask_dilation_mm: float) -> None: """ self.mask_dilation_mm = mask_dilation_mm + def set_distancemap_squared_max(self, distancemap_squared_max: float) -> None: + """Set the saturation radius of the labelmap-to-labelmap distance maps. + + The radius fixes those images' intensity distribution, so it has to + match the value the ICON weights in use were finetuned at -- see + ``tutorials/parameters_lung_ct_dirlab.py`` and + ``tutorials/parameters_heart_ct_kcl.py``, which carry one value per + organ. Left unset, it is derived from ``mask_dilation_mm``. + + Args: + distancemap_squared_max: Maximum squared distance, in squared + millimeters, the distance maps are normalized against. + """ + self.distancemap_squared_max = distancemap_squared_max + + def _distancemap_squared_max(self) -> float: + """Return the configured saturation radius, or one sized to the mask.""" + if self.distancemap_squared_max is not None: + return self.distancemap_squared_max + return (1.25 * self.mask_dilation_mm) ** 2 + def set_labelmap_to_labelmap_icon_weights_path(self, weights_path: str) -> None: """Set a finetuned ICON checkpoint for the labelmap-to-labelmap stage. @@ -719,7 +741,7 @@ def register_labelmap_to_labelmap(self) -> Optional[dict]: fixed_model=self.patient_model_surface, reference_image=padded_patient_image, mask_dilation_mm=self.mask_dilation_mm, - distance_squared_max=(1.25 * self.mask_dilation_mm) ** 2, + distance_squared_max=self._distancemap_squared_max(), ) if self.l2l_icon_weights_path is not None: labelmap_registrar.set_icon_weights_path(self.l2l_icon_weights_path) diff --git a/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py b/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py index 3d1ed2db..b0c28036 100644 --- a/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py +++ b/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py @@ -55,8 +55,7 @@ class WorkflowReconstructHighres4DCT(PhysioTwin4DBase): ``registration_method`` accepts a pre-configured :class:`RegisterImagesBase` instance. Configure backend-specific parameters (iteration counts, etc.) on the instance before passing it - in. Defaults to a new :class:`RegisterImagesGreedyICON` (Greedy followed - by ICON refinement) when omitted. + in. Defaults to a new :class:`RegisterImagesGreedy` when omitted. Attributes: time_series_images (list[itk.Image]): Ordered list of time-series images diff --git a/tests/test_workflow_convert_image_to_usd.py b/tests/test_workflow_convert_image_to_usd.py index d7ede866..78f7c83e 100644 --- a/tests/test_workflow_convert_image_to_usd.py +++ b/tests/test_workflow_convert_image_to_usd.py @@ -12,6 +12,7 @@ from pxr import Usd, UsdGeom from physiotwin4d.register_images_base import RegisterImagesBase +from physiotwin4d.register_images_greedy import RegisterImagesGreedy from physiotwin4d.register_images_icon import RegisterImagesICON from physiotwin4d.segment_chest_total_segmentator_with_contrast import ( SegmentChestTotalSegmentatorWithContrast, @@ -27,7 +28,7 @@ def _small_image() -> itk.Image: def test_default_segmentation_and_registration_methods(tmp_path: Path) -> None: """Omitting segmentation_method/registration_method defaults to SegmentChestTotalSegmentatorWithContrast (contrast_threshold=500) and - RegisterImagesICON, matching this workflow's documented defaults.""" + RegisterImagesGreedy, matching this workflow's documented defaults.""" reference_image = _small_image() workflow = WorkflowConvertImageToUSD( time_series_images=[reference_image], @@ -39,7 +40,7 @@ def test_default_segmentation_and_registration_methods(tmp_path: Path) -> None: assert isinstance(workflow.segmenter, SegmentChestTotalSegmentatorWithContrast) assert workflow.segmenter.contrast_threshold == 500 - assert isinstance(workflow.registrar, RegisterImagesICON) + assert isinstance(workflow.registrar, RegisterImagesGreedy) def test_segmentation_method_rejects_wrong_type(tmp_path: Path) -> None: @@ -117,8 +118,8 @@ def test_workflow_convert_image_to_usd_default_operation( assert isinstance(workflow.segmenter, SegmentChestTotalSegmentatorWithContrast) assert workflow.segmenter.contrast_threshold == 500 - assert isinstance(workflow.registrar, RegisterImagesICON) - workflow.registrar.set_number_of_iterations(2) + assert isinstance(workflow.registrar, RegisterImagesGreedy) + workflow.registrar.set_number_of_iterations([2]) result_filenames = workflow.process() diff --git a/tests/test_workflow_reconstruct_highres_4d_ct.py b/tests/test_workflow_reconstruct_highres_4d_ct.py index 1c15c745..fb7aeacd 100644 --- a/tests/test_workflow_reconstruct_highres_4d_ct.py +++ b/tests/test_workflow_reconstruct_highres_4d_ct.py @@ -8,7 +8,7 @@ import pytest from physiotwin4d.register_images_base import RegisterImagesBase -from physiotwin4d.register_images_greedy_icon import RegisterImagesGreedyICON +from physiotwin4d.register_images_greedy import RegisterImagesGreedy from physiotwin4d.register_images_icon import RegisterImagesICON from physiotwin4d.workflow_reconstruct_highres_4d_ct import ( WorkflowReconstructHighres4DCT, @@ -20,14 +20,13 @@ def _small_image() -> itk.Image: return itk.image_from_array(np.zeros((3, 3, 3), dtype=np.float32)) -def test_default_registration_method_is_greedy_icon() -> None: - """Omitting registration_method defaults to RegisterImagesGreedyICON, - matching this workflow's historical 'Greedy_ICON' string default.""" +def test_default_registration_method_is_greedy() -> None: + """Omitting registration_method defaults to RegisterImagesGreedy.""" workflow = WorkflowReconstructHighres4DCT( time_series_images=[_small_image(), _small_image()], reference_image=_small_image(), ) - assert isinstance(workflow.registrar.registrar, RegisterImagesGreedyICON) + assert isinstance(workflow.registrar.registrar, RegisterImagesGreedy) def test_registration_method_rejects_wrong_type() -> None: diff --git a/tutorials/README.md b/tutorials/README.md index 02c37c7f..71cba42b 100644 --- a/tutorials/README.md +++ b/tutorials/README.md @@ -28,7 +28,8 @@ current working directory. | 1 | [tutorial_01_heart_gated_ct_to_usd.py](tutorial_01_heart_gated_ct_to_usd.py) | `WorkflowConvertImageToUSD` | Slicer-Heart-CT (prepare first) | | 1 | [tutorial_01_lung_gated_ct_to_usd.py](tutorial_01_lung_gated_ct_to_usd.py) | `WorkflowConvertImageToUSD` | Lung gated 4D CT (prepare first) | | 2 | [tutorial_02_lung_finetune_icon.py](tutorial_02_lung_finetune_icon.py) | `WorkflowFinetuneICONRegistration` | DirLab-4DCT (manual) | -| 2 | [distancemap variant](tutorial_02_lung_distancemap_finetune_icon.py) | `WorkflowFinetuneICONRegistration` on distance maps | DirLab-4DCT (manual) | +| 2 | [lung distancemap variant](tutorial_02_lung_distancemap_finetune_icon.py) | `WorkflowFinetuneICONRegistration` on distance maps | DirLab-4DCT (manual) | +| 2 | [heart distancemap variant](tutorial_02_duke_heart_distancemap_finetune_icon.py) | `WorkflowFinetuneICONRegistration` on distance maps | Duke-Heart-4DLabelmaps (not yet available) | | 3 | [tutorial_03_heart_reconstruct_highres_4d_ct.py](tutorial_03_heart_reconstruct_highres_4d_ct.py) | `WorkflowReconstructHighres4DCT` | Slicer-Heart-CT (prepare first) | | 3 | [tutorial_03_lung_reconstruct_highres_4d_ct.py](tutorial_03_lung_reconstruct_highres_4d_ct.py) | `WorkflowReconstructHighres4DCT` | DirLab-4DCT (manual) | | 4 | [tutorial_04_heart_ct_to_vtk.py](tutorial_04_heart_ct_to_vtk.py) | `WorkflowConvertImageToVTK` | Slicer-Heart-CT (prepare first) | diff --git a/tutorials/parameters_heart_ct_kcl.py b/tutorials/parameters_heart_ct_kcl.py new file mode 100644 index 00000000..4d4bfa14 --- /dev/null +++ b/tutorials/parameters_heart_ct_kcl.py @@ -0,0 +1,89 @@ +"""Shared parameters for the heart CT tutorials. + +Mirrors :mod:`parameters_lung_ct_dirlab` for the heart use cases, and carries +different values: the heart is registered with a much tighter mask than the +lungs, so its distance maps saturate over a correspondingly shorter radius. +That is why the heart has its own distance-map finetuning tutorial rather than +reusing the lung one's weights -- the two organs' distance maps do not share an +intensity distribution. + +Paths stay out of this module: each tutorial owns its own inputs and outputs. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Optional + +from physiotwin4d import SegmentAnatomyBase, SegmentChestTotalSegmentator + + +@dataclass(frozen=True) +class ParametersHeartCTKCL: + """Settings shared by the heart tutorials. + + Attributes: + mask_dilation_mm: Dilation of the binary registration masks, in + millimeters. Tighter than the lungs': the heart is a compact organ + whose neighbours are not part of the model. + distancemap_squared_max: Saturation radius of every heart distance map, + in squared millimeters. Fixes their intensity distribution, so the + finetuning tutorial and every tutorial that registers heart distance + maps must use this one value. + number_of_pca_components: PCA components retained when building the + heart statistical model, and used when fitting it to a patient. + number_of_pca_components_test: Same, under ``TestTools.running_as_test``. + number_of_iterations_greedy: Greedy coarse-to-fine iteration schedule. + number_of_iterations_greedy_test: Same, under + ``TestTools.running_as_test``. + segmenter_class: Segmenter every heart tutorial instantiates, so the + surfaces they compare share a definition of "heart". + anatomy_group: Anatomy group name that segmenter registers for the heart. + interior_object_ids_totalsegmentator: Chamber labels in a + TotalSegmentator labelmap. The chambers are interior to the + myocardium, so a distance map must not measure to them. + interior_object_ids_simpleware: The same chambers in a Simpleware + ASCardio labelmap, which the Duke-Heart-4DLabelmaps data was + segmented with. Which list applies is a property of the data a + tutorial reads, not of any one segmenter class, so both live here. + """ + + mask_dilation_mm: float = 10.0 + distancemap_squared_max: float = (1.25 * 10.0) ** 2 + + number_of_pca_components: int = 10 + number_of_pca_components_test: int = 5 + + number_of_iterations_greedy: list[int] = field( + default_factory=lambda: [30, 15, 7, 3] + ) + number_of_iterations_greedy_test: list[int] = field(default_factory=lambda: [1, 0]) + + segmenter_class: type[SegmentAnatomyBase] = SegmentChestTotalSegmentator + anatomy_group: str = "heart" + interior_object_ids_totalsegmentator: Optional[list[int]] = field( + default_factory=lambda: [141, 142, 143, 144] + ) + interior_object_ids_simpleware: list[int] = field( + default_factory=lambda: [1, 2, 3, 4] + ) + + def pca_components(self, test_mode: bool) -> int: + """Return the PCA component count for this run mode.""" + return ( + self.number_of_pca_components_test + if test_mode + else (self.number_of_pca_components) + ) + + def greedy_iterations(self, test_mode: bool) -> list[int]: + """Return the Greedy iteration schedule for this run mode.""" + return list( + self.number_of_iterations_greedy_test + if test_mode + else self.number_of_iterations_greedy + ) + + +#: The single instance every heart tutorial imports. +HEART_CT_KCL = ParametersHeartCTKCL() diff --git a/tutorials/parameters_lung_ct_dirlab.py b/tutorials/parameters_lung_ct_dirlab.py new file mode 100644 index 00000000..08cbb684 --- /dev/null +++ b/tutorials/parameters_lung_ct_dirlab.py @@ -0,0 +1,78 @@ +"""Shared parameters for the DIR-Lab 4D CT lung tutorials. + +Every lung tutorial reads its settings from :data:`LUNG_CT_DIRLAB` so the +distance maps one tutorial finetunes ICON on are rasterized exactly the way the +tutorials that later register them rasterize theirs. A saturation radius or a +dilation that drifts between two of these scripts silently trains on one image +distribution and infers on another. + +Paths stay out of this module: each tutorial owns its own inputs and outputs. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from physiotwin4d import SegmentAnatomyBase, SegmentNVSegmentCTMRI + + +@dataclass(frozen=True) +class ParametersLungCTDirLab: + """Settings shared by the DIR-Lab lung tutorials. + + Attributes: + mask_dilation_mm: Dilation of the binary registration masks, in + millimeters. Also sets how far outside the lung surface the + registration is allowed to look. + distancemap_squared_max: Saturation radius of every lung distance map, + in squared millimeters. Fixes their intensity distribution, so the + finetuning tutorial and every tutorial that registers lung distance + maps must use this one value. + number_of_pca_components: PCA components retained when building the + lung statistical model, and used when fitting it to a patient. + number_of_pca_components_test: Same, under ``TestTools.running_as_test``. + number_of_iterations_greedy: Greedy coarse-to-fine iteration schedule. + number_of_iterations_greedy_test: Same, under + ``TestTools.running_as_test``. + segmenter_class: Segmenter every lung tutorial instantiates, so the + surfaces they compare share a definition of "lung". + anatomy_group: Anatomy group name that segmenter registers for lungs. + + There is no interior-structure list here, the counterpart of the heart's + chamber ids: the lung labels are the lobes, and every one of them is on the + surface a distance map is measured to. + """ + + mask_dilation_mm: float = 40.0 + distancemap_squared_max: float = (1.25 * 40.0) ** 2 + + number_of_pca_components: int = 6 + number_of_pca_components_test: int = 5 + + number_of_iterations_greedy: list[int] = field( + default_factory=lambda: [30, 15, 7, 3] + ) + number_of_iterations_greedy_test: list[int] = field(default_factory=lambda: [1, 0]) + + segmenter_class: type[SegmentAnatomyBase] = SegmentNVSegmentCTMRI + anatomy_group: str = "lung" + + def pca_components(self, test_mode: bool) -> int: + """Return the PCA component count for this run mode.""" + return ( + self.number_of_pca_components_test + if test_mode + else (self.number_of_pca_components) + ) + + def greedy_iterations(self, test_mode: bool) -> list[int]: + """Return the Greedy iteration schedule for this run mode.""" + return list( + self.number_of_iterations_greedy_test + if test_mode + else self.number_of_iterations_greedy + ) + + +#: The single instance every lung tutorial imports. +LUNG_CT_DIRLAB = ParametersLungCTDirLab() diff --git a/tutorials/tutorial_01_heart_gated_ct_to_usd.py b/tutorials/tutorial_01_heart_gated_ct_to_usd.py index f0e2b4fa..dcaa4b72 100644 --- a/tutorials/tutorial_01_heart_gated_ct_to_usd.py +++ b/tutorials/tutorial_01_heart_gated_ct_to_usd.py @@ -31,15 +31,12 @@ Strengths --------- - Single call (``WorkflowConvertImageToUSD.process()``) runs the full pipeline. -- Supports both GPU-accelerated ICON registration and CPU-capable Greedy registration. +- Registers on the CPU with ``RegisterImagesGreedy``; no GPU needed for this stage. - Automatically detects contrast enhancement and adjusts segmentation thresholds. - Output is Omniverse-ready with anatomical materials (USDAnatomyTools). Weaknesses / Limitations ------------------------ -- Requires a GPU for ICON registration (``registration_method=RegisterImagesICON()``); - use ``registration_method=RegisterImagesGreedy()`` for CPU-only environments - (about 10x slower). - Segmentation quality depends on TotalSegmentator's training distribution; unusual pathologies or pediatric anatomy may degrade results. - Large 4D datasets (>20 phases, high resolution) can require 32 GB+ RAM. @@ -51,7 +48,7 @@ contour extraction -> USD export. - SegmentChestTotalSegmentator (segment_chest_total_segmentator.py): Deep-learning segmentation of 117 anatomical structures (used internally). -- RegisterImagesICON / RegisterImagesANTS (register_images_icon.py / _ants.py): +- RegisterImagesGreedy (register_images_greedy.py): Frame-to-frame image registration (used internally). - ContourTools (contour_tools.py): Extracts and transforms surface meshes from segmentation masks (used internally). @@ -76,7 +73,7 @@ import itk from physiotwin4d import ( - RegisterImagesICON, + RegisterImagesGreedy, SegmentChestTotalSegmentatorWithContrast, TestTools, WorkflowConvertImageToUSD, @@ -101,17 +98,17 @@ test_mode = TestTools.running_as_test() if test_mode: data_dir = repo_root / "data" / "test" / "slicer_heart_small" - number_of_registration_iterations = 1 + number_of_iterations_greedy = [1, 0] frame_files = sorted(data_dir.glob("slice_???.mha"))[0:2] else: data_dir = repo_root / "data" / "Slicer-Heart-CT" - number_of_registration_iterations = 10 + number_of_iterations_greedy = [30, 15, 7, 3] frame_files = sorted(data_dir.glob("slice_???.mha")) log_level = logging.INFO - registration_method = RegisterImagesICON(log_level=log_level) - registration_method.set_number_of_iterations(number_of_registration_iterations) + registration_method = RegisterImagesGreedy(log_level=log_level) + registration_method.set_number_of_iterations(number_of_iterations_greedy) segmentation_method = SegmentChestTotalSegmentatorWithContrast(log_level=log_level) segmentation_method.set_has_academic_license(True) diff --git a/tutorials/tutorial_01_lung_gated_ct_to_usd.py b/tutorials/tutorial_01_lung_gated_ct_to_usd.py index 18009538..d34e7f3c 100644 --- a/tutorials/tutorial_01_lung_gated_ct_to_usd.py +++ b/tutorials/tutorial_01_lung_gated_ct_to_usd.py @@ -30,14 +30,11 @@ Strengths --------- - Single call (``WorkflowConvertImageToUSD.process()``) runs the full pipeline. -- Supports both GPU-accelerated ICON registration and CPU-capable Greedy registration. +- Registers on the CPU with ``RegisterImagesGreedy``; no GPU needed for this stage. - Output is Omniverse-ready with anatomical materials (USDAnatomyTools). Weaknesses / Limitations ------------------------ -- Requires a GPU for ICON registration (``registration_method=RegisterImagesICON()``); - use ``registration_method=RegisterImagesGreedy()`` for CPU-only environments - (about 10x slower). - Segmentation quality depends on TotalSegmentator's training distribution; unusual pathologies or pediatric anatomy may degrade results. - Large 4D datasets (>20 phases, high resolution) can require 32 GB+ RAM. @@ -49,7 +46,7 @@ contour extraction -> USD export. - SegmentChestTotalSegmentator (segment_chest_total_segmentator.py): Deep-learning segmentation of 117 anatomical structures (used internally). -- RegisterImagesICON (register_images_icon.py): +- RegisterImagesGreedy (register_images_greedy.py): Frame-to-frame image registration (used internally). - ContourTools (contour_tools.py): Extracts and transforms surface meshes from segmentation masks (used internally). @@ -74,7 +71,7 @@ import itk from physiotwin4d import ( - RegisterImagesICON, + RegisterImagesGreedy, SegmentChestTotalSegmentator, TestTools, WorkflowConvertImageToUSD, @@ -102,16 +99,16 @@ # data/DirLab-4DCT/fix_downloaded_data.py. test_mode = TestTools.running_as_test() if test_mode: - number_of_registration_iterations = 1 + number_of_iterations_greedy = [1, 0] frame_files = sorted(data_dir.glob("Case1Pack_T??.mha"))[0:2] else: - number_of_registration_iterations = 10 + number_of_iterations_greedy = [30, 15, 7, 3] frame_files = sorted(data_dir.glob("Case1Pack_T??.mha")) log_level = logging.INFO - registration_method = RegisterImagesICON(log_level=log_level) - registration_method.set_number_of_iterations(number_of_registration_iterations) + registration_method = RegisterImagesGreedy(log_level=log_level) + registration_method.set_number_of_iterations(number_of_iterations_greedy) segmentation_method = SegmentChestTotalSegmentator(log_level=log_level) segmentation_method.set_has_academic_license(True) diff --git a/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py b/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py new file mode 100644 index 00000000..34658b4a --- /dev/null +++ b/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py @@ -0,0 +1,471 @@ +""" +Tutorial 2 (Heart): Finetune uniGradICON on heart distance maps + +Purpose +------- +The heart counterpart of ``tutorial_02_lung_distancemap_finetune_icon.py``. +Both finetune uniGradICON on *distance maps* rather than on CT intensities, +because that is what ``RegisterModelsDistanceMaps`` -- the labelmap-to-labelmap +stage of ``WorkflowFitStatisticalModelToPatient`` -- actually registers. Stock +uniGradICON weights are out of distribution for those images. + +The heart needs its own run rather than reusing the lung weights. A distance +map's appearance is set by the radius it saturates at, and the heart is +registered with a far tighter mask than the lungs +(``ParametersHeartCTKCL.mask_dilation_mm`` versus +``ParametersLungCTDirLab.mask_dilation_mm``), so the two organs' distance maps +do not share an intensity distribution. Every value that fixes that +distribution comes from ``parameters_heart_ct_kcl.py``, so this run trains on +exactly what Tutorial 7 (heart) later infers on. + +Duke-Heart-4DLabelmaps ships segmented labelmaps, not CT, which is all this +tutorial needs. Each labelmap serves twice: with the chamber-interior labels +dropped it yields the surface a distance map is measured to, and whole it +supplies uniGradICON's Dice loss. Those labelmaps were produced by +``SegmentHeartSimplewareTrimmedBranches``, whose ventricle and atrium labels +cover the cavities rather than the walls, so measuring distance to them would +measure distance to the inside of the heart; the ids to drop come from +``parameters_heart_ct_kcl.py`` alongside every other value this data needs. + +The first patient is held out and two of its gated phases are registered three +ways -- Greedy on the distance maps, stock ICON, and the finetuned ICON -- and +scored by target registration error over the anatomical landmarks Slicer markups +files supply for every frame, plus per-class Dice. + +Data Required +------------- +``data/Duke-Heart-4DLabelmaps/pm????/`` with, per gated frame: +``*_labelmap.nii.gz`` (multi-label) and ``*_landmark.mrk.json`` (Slicer markups, +LPS). +""" + +# Imports +from __future__ import annotations + +import csv +import json +import logging +import shutil +import time +from pathlib import Path +from typing import Any, Optional + +import itk +import numpy as np +from parameters_heart_ct_kcl import HEART_CT_KCL + +from physiotwin4d import ( + ContourTools, + PhysioTwin4DBase, + RegisterImagesBase, + RegisterImagesGreedy, + RegisterImagesICON, + TestTools, + TransformTools, + WorkflowFinetuneICONRegistration, +) + +# Only run if this script is not imported as a module + +# unigradicon finetuning is launched as a subprocess and torch spawns worker +# processes; on Windows the spawn start method re-imports this script in each +# child, so all top-level work stays under the __name__ == "__main__" guard. +if __name__ == "__main__": + # Data directory specification + repo_root = Path(__file__).resolve().parent.parent + tutorials_dir = Path(__file__).resolve().parent + + class_name = "tutorial_02_duke_heart_distancemap_finetune_icon" + + output_dir = tutorials_dir / "output" / "tutorial_02_heart_distancemap" + # Rasterized distance maps, cached so re-runs skip the rasterization. + derived_dir = output_dir / "distance_maps" + baselines_dir = repo_root / "tests" / "baselines" + + # The workflow writes its dataset JSON, YAML config, and checkpoint tree + # under ``weights_dir / finetune_name``. + weights_dir = tutorials_dir / "network_weights" + finetune_name = "icon_duke_heart_distancemap" + + # Distance-map normalization, shared with every heart tutorial that later + # registers these maps, so this run trains on the same image distribution + # they infer on. + distance_squared_max = HEART_CT_KCL.distancemap_squared_max + + # Set True to finetune from scratch. That deletes experiment_dir below, + # including any checkpoint a previous run left there. + run_finetuning = True + + test_mode = TestTools.running_as_test() + if test_mode: + data_dir = repo_root / "data" / "test" / "Duke-Heart-4DLabelmaps" + number_of_iterations_icon: Optional[int] = 1 + epochs = 1 + else: + data_dir = repo_root / "data" / "Duke-Heart-4DLabelmaps" + number_of_iterations_icon = 10 + epochs = 100 + number_of_iterations_greedy = HEART_CT_KCL.greedy_iterations(test_mode) + + log_level = logging.INFO + reporter = PhysioTwin4DBase(class_name=class_name, log_level=log_level) + + derived_dir.mkdir(parents=True, exist_ok=True) + + contour_tools = ContourTools(log_level=log_level) + transform_tools = TransformTools() + + # This data was segmented with Simpleware ASCardio, so the chamber labels + # to drop are that labelmap's, not TotalSegmentator's. + interior_object_ids = HEART_CT_KCL.interior_object_ids_simpleware + + # Cohort discovery + case_dirs = sorted(path for path in data_dir.glob("pm*") if path.is_dir()) + if len(case_dirs) < 2: + raise FileNotFoundError( + f"Need at least 2 Duke heart cases under {data_dir}; found " + f"{len(case_dirs)}.\nSee data/README.md for download instructions." + ) + + def frames_for_case(case_dir: Path) -> list[Path]: + """Return one path stem per gated frame, as its multi-label labelmap.""" + return sorted(case_dir.glob("*_labelmap.nii.gz")) + + def companion(labelmap_file: Path, suffix: str) -> Path: + """Return the file that shares *labelmap_file*'s frame stem.""" + stem = labelmap_file.name[: -len("_labelmap.nii.gz")] + return labelmap_file.parent / f"{stem}{suffix}" + + def distance_map_for(labelmap_file: Path) -> Path: + """Rasterize one frame's heart distance map, caching it under derived_dir. + + Mirrors ``RegisterModelsDistanceMaps._create_masks_from_models`` so the + finetuning inputs match what that class feeds ICON at inference: a + signed squared distance to the heart surface, normalized to [-1, 1] by + ``distance_squared_max``, then multiplied by 1000 to fill the + [-1000, 1000] window uniGradICON's CT preprocessing expects. + """ + stem = labelmap_file.name[: -len("_labelmap.nii.gz")] + distance_map_file = derived_dir / f"{stem}_distance_map.mha" + if distance_map_file.exists(): + return distance_map_file + + reporter.log_info("Rasterizing distance map for %s", stem) + labelmap_image = itk.imread(str(labelmap_file)) + labels = itk.GetArrayViewFromImage(labelmap_image) + # The chamber labels cover the cavities rather than their walls, so a + # distance measured to them would be a distance to the inside of the + # heart. Dropping them leaves the myocardium and vessels, whose + # boundary is the surface RegisterModelsDistanceMaps measures to at + # inference. + heart_mask = itk.GetImageFromArray( + np.where((labels > 0) & ~np.isin(labels, interior_object_ids), 1, 0).astype( + np.uint8 + ) + ) + heart_mask.CopyInformation(labelmap_image) + surface = contour_tools.extract_contours(heart_mask) + distance_map = contour_tools.create_distance_map( + surface, + labelmap_image, + squared_distance=True, + negative_inside=True, + zero_inside=False, + norm_to_max_distance=distance_squared_max, + ) + itk.GetArrayViewFromImage(distance_map)[...] *= 1000 + itk.imwrite(distance_map, str(distance_map_file), compression=True) + return distance_map_file + + # Held-out patient: the first case, excluded from finetuning entirely. + held_out_dir = case_dirs[0] + training_dirs = case_dirs[1:] + + held_out_frames = frames_for_case(held_out_dir) + if len(held_out_frames) < 2: + raise FileNotFoundError( + f"Held-out case {held_out_dir.name} has {len(held_out_frames)} frames; " + "at least 2 are needed to form an evaluation pair." + ) + # The 30% and 70% gated phases are the most separated, so they are the + # hardest pair in the case and the most informative to score. + moving_labelmap_file = held_out_frames[int(len(held_out_frames) * 0.3)] + fixed_labelmap_file = held_out_frames[int(len(held_out_frames) * 0.7)] + + subject_distance_map_files: list[list[str]] = [] + subject_labelmap_files: list[list[Optional[str]]] = [] + subject_ids: list[str] = [] + for case_dir in training_dirs: + frames = frames_for_case(case_dir) + if len(frames) < 2: + reporter.log_warning( + "Case %s has %d frame(s); skipping (paired training needs 2+)", + case_dir.name, + len(frames), + ) + continue + subject_ids.append(case_dir.name) + subject_distance_map_files.append( + [str(distance_map_for(frame)) for frame in frames] + ) + subject_labelmap_files.append([str(frame) for frame in frames]) + reporter.log_info( + "Finetuning cohort: %d cases, %d frames (held out %s)", + len(subject_ids), + sum(len(files) for files in subject_distance_map_files), + held_out_dir.name, + ) + + # Always finetune from scratch. uniGradICON refuses to overwrite an + # existing experiment directory: it appends "-N" to the name instead + # (``icon_duke_heart_distancemap_model-5``, ...), while + # expected_weights_path() keeps pointing at the original, never-written + # path. Deleting the tree up front keeps the two in agreement. + experiment_dir = weights_dir / finetune_name + if run_finetuning: + if experiment_dir.exists(): + reporter.log_info( + "Removing previous finetuning outputs: %s", experiment_dir + ) + shutil.rmtree(experiment_dir) + + # The labelmaps are supplied, so uniGradICON's Dice loss stays enabled + # at its default weight. They hold only the heart's own structures, + # which keeps the one-hot encoding the Dice term builds small. + # + # lncc_sigma matches the sigma RegisterImagesICON uses at inference, so + # finetuning optimizes the similarity this comparison scores. The + # distance maps are already scaled into [-1000, 1000], so the default CT + # window passes them through unclipped. + workflow = WorkflowFinetuneICONRegistration( + subject_image_files=subject_distance_map_files, + output_dir=weights_dir, + finetune_name=finetune_name, + subject_ids=subject_ids, + subject_labelmap_files=subject_labelmap_files, + epochs=epochs, + lncc_sigma=5, + log_level=log_level, + ) + weights_path = workflow.process() + else: + weights_path = ( + experiment_dir + / f"{finetune_name}_model" + / "checkpoints" + / "network_weights_final.trch" + ) + # Checked here rather than at the first set_weights_path() call, which + # only happens after the greedy and stock-ICON rows have already run. + if not weights_path.exists(): + raise FileNotFoundError( + f"run_finetuning is False but no checkpoint at {weights_path}. " + "Set run_finetuning = True to create one." + ) + + # Registration comparison on the held-out patient + fixed_distance_map = itk.imread( + str(distance_map_for(fixed_labelmap_file)), pixel_type=itk.F + ) + moving_distance_map = itk.imread( + str(distance_map_for(moving_labelmap_file)), pixel_type=itk.F + ) + fixed_labelmap = itk.imread(str(fixed_labelmap_file)) + moving_labelmap = itk.imread(str(moving_labelmap_file)) + + def read_landmarks(labelmap_file: Path) -> dict[str, np.ndarray]: + """Read a frame's Slicer markups file as ``{label: LPS point}``. + + The markups files declare ``coordinateSystem: LPS``, the frame this + project works in, so the control points are used as written. + """ + landmark_file = companion(labelmap_file, "_landmark.mrk.json") + with landmark_file.open(encoding="utf-8") as f: + markups = json.load(f)["markups"] + return { + point["label"]: np.asarray(point["position"], dtype=np.float64) + for markup in markups + for point in markup["controlPoints"] + } + + fixed_landmarks_by_name = read_landmarks(fixed_labelmap_file) + moving_landmarks_by_name = read_landmarks(moving_labelmap_file) + shared_landmarks = sorted( + set(fixed_landmarks_by_name) & set(moving_landmarks_by_name) + ) + if not shared_landmarks: + raise ValueError( + f"No landmark names shared between {fixed_labelmap_file.name} and " + f"{moving_labelmap_file.name}." + ) + fixed_landmarks = np.array( + [fixed_landmarks_by_name[name] for name in shared_landmarks] + ) + moving_landmarks = np.array( + [moving_landmarks_by_name[name] for name in shared_landmarks] + ) + reporter.log_info( + "Scoring %d shared landmarks on held-out case %s", + len(shared_landmarks), + held_out_dir.name, + ) + + def landmark_metrics(errors_mm: np.ndarray) -> dict[str, Any]: + """Summarize per-landmark target registration errors, in millimeters.""" + return { + "tre_mean": float(errors_mm.mean()), + "tre_std": float(errors_mm.std()), + "tre_max": float(errors_mm.max()), + } + + def landmark_errors(transform: itk.Transform) -> np.ndarray: + """Distance from each mapped fixed landmark to its moving counterpart. + + ``forward_transform`` is the resampling transform: it maps points on the + fixed grid back into moving space, which is the direction the landmark + correspondences are defined in. + """ + mapped = np.array( + [transform.TransformPoint(tuple(point)) for point in fixed_landmarks] + ) + return np.asarray(np.linalg.norm(mapped - moving_landmarks, axis=1)) + + fixed_labels = itk.array_from_image(fixed_labelmap) + + def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: + """Per-class Dice summary against the fixed labelmap. + + Classes are the union of the two labelmaps' non-zero ids, so a class + found in only one of them scores 0 rather than being dropped. + """ + labels = itk.array_from_image(labelmap) + classes = np.union1d(np.unique(fixed_labels), np.unique(labels)) + classes = classes[classes != 0] + dice = np.array( + [ + 2.0 + * np.count_nonzero((fixed_labels == c) & (labels == c)) + / (np.count_nonzero(fixed_labels == c) + np.count_nonzero(labels == c)) + for c in classes + ] + ) + return { + "n_classes": int(dice.size), + "dice_mean": float(dice.mean()), + "dice_min": float(dice.min()), + "dice_max": float(dice.max()), + "mislabeled_voxels": int(np.count_nonzero(fixed_labels != labels)), + } + + # Reference row: the moving labelmap on the fixed grid, unregistered. + unregistered_labelmap = itk.resample_image_filter( + moving_labelmap, + Interpolator=itk.NearestNeighborInterpolateImageFunction.New(moving_labelmap), + ReferenceImage=fixed_labelmap, + UseReferenceImage=True, + ) + labelmaps: dict[str, itk.Image] = {"unregistered": unregistered_labelmap} + rows: list[dict[str, Any]] = [ + { + "method": "unregistered", + "weights": "-", + "registration_time_s": None, + **landmark_metrics( + np.linalg.norm(fixed_landmarks - moving_landmarks, axis=1) + ), + **overlap_metrics(unregistered_labelmap), + } + ] + for method_name, method_weights in ( + ("greedy", None), + ("icon_stock", None), + ("icon_finetuned", weights_path), + ): + registrar: RegisterImagesBase + if method_name == "greedy": + registrar = RegisterImagesGreedy(log_level=log_level) + registrar.set_transform_type("Deformable") + # CC is what RegisterModelsDistanceMaps uses on distance maps. + registrar.set_metric("CC") + registrar.set_number_of_iterations(number_of_iterations_greedy) + else: + registrar = RegisterImagesICON(log_level=log_level) + # None, not 0: icon_registration rejects 0 and takes None to mean + # "no test-time finetuning steps". + registrar.set_number_of_iterations(number_of_iterations_icon) + # Mass preservation models CT density; a distance map carries no + # mass, so it stays off here. + registrar.set_mass_preservation(False) + if method_weights is not None: + registrar.set_weights_path(str(method_weights)) + registrar.set_modality("ct") + registrar.set_fixed_image(fixed_distance_map) + + start_time = time.perf_counter() + result = registrar.register(moving_distance_map) + elapsed_s = time.perf_counter() - start_time + + labelmaps[method_name] = transform_tools.transform_image( + moving_labelmap, + result["forward_transform"], + fixed_labelmap, + interpolation_method="nearest", + ) + rows.append( + { + "method": method_name, + "weights": str(method_weights) if method_weights else "-", + "registration_time_s": elapsed_s, + **landmark_metrics(landmark_errors(result["forward_transform"])), + **overlap_metrics(labelmaps[method_name]), + } + ) + + # Result saving + output_dir.mkdir(parents=True, exist_ok=True) + for method_name, labelmap in labelmaps.items(): + itk.imwrite( + labelmap, + str(output_dir / f"labelmap_{method_name}.mha"), + compression=True, + ) + + summary_file = output_dir / "registration_summary.csv" + with summary_file.open("w", newline="", encoding="utf-8") as fh: + writer = csv.DictWriter(fh, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + for row in rows: + reporter.log_info( + "%-16s TRE mean %s mm, Dice mean %.4f", + row["method"], + f"{row['tre_mean']:.4f}", + row["dice_mean"], + ) + + # Testing + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + baselines_dir=baselines_dir, + log_level=log_level, + ) + + screenshots: list[Path] = [ + tt.save_screenshot_image_slice( + fixed_distance_map, + "fixed_distance_map.png", + axis=0, + slice_fraction=0.5, + colormap="gray", + ) + ] + + tutorial_results = { + "weights_path": weights_path, + "registration_metrics": rows, + "labelmaps": labelmaps, + "summary_file": summary_file, + "screenshots": screenshots, + } diff --git a/tutorials/tutorial_02_lung_distancemap_finetune_icon.py b/tutorials/tutorial_02_lung_distancemap_finetune_icon.py index 1d754140..4659790f 100644 --- a/tutorials/tutorial_02_lung_distancemap_finetune_icon.py +++ b/tutorials/tutorial_02_lung_distancemap_finetune_icon.py @@ -79,6 +79,8 @@ import numpy as np import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB + from physiotwin4d import ( ContourTools, PhysioTwin4DBase, @@ -117,14 +119,10 @@ weights_dir = tutorials_dir / "network_weights" finetune_name = "icon_dirlab_4dct_distancemap" - # Distance-map normalization. WorkflowFitStatisticalModelToPatient passes - # (1.25 * mask_dilation_mm) ** 2 as RegisterModelsDistanceMaps' - # distance_squared_max, so the value below fixes the saturation radius of - # every distance map the finetuned weights will ever see. The lung fitting - # tutorials set the same mask_dilation_mm; changing one without the other - # trains on a different image distribution than it infers on. - mask_dilation_mm = 40.0 - distance_squared_max = (1.25 * mask_dilation_mm) ** 2 + # Distance-map normalization, shared with every lung tutorial that later + # registers these maps, so this run trains on the same image distribution + # they infer on. + distance_squared_max = LUNG_CT_DIRLAB.distancemap_squared_max run_finetuning = True @@ -135,14 +133,13 @@ test_mode = TestTools.running_as_test() if test_mode: data_dir = repo_root / "data" / "test" / "DirLab-4DCT" - number_of_iterations_greedy: Optional[list[int]] = [1, 0] - number_of_iterations_icon = 1 + number_of_iterations_icon: Optional[int] = 1 epochs = 1 else: data_dir = repo_root / "data" / "DirLab-4DCT" - number_of_iterations_greedy = [60, 30, 20] number_of_iterations_icon = 10 epochs = 200 + number_of_iterations_greedy = LUNG_CT_DIRLAB.greedy_iterations(test_mode) log_level = logging.INFO reporter = PhysioTwin4DBase(class_name=class_name, log_level=log_level) @@ -169,10 +166,20 @@ ) # Segmentation and distance-map generation + segmenter = SegmentNVSegmentCTMRI(log_level=log_level) segmentation_workflow = WorkflowConvertImageToVTK( - segmentation_method=SegmentNVSegmentCTMRI(log_level=log_level), + segmentation_method=segmenter, log_level=log_level, ) + # WorkflowConvertImageToVTK returns the segmenter's whole-body labelmap + # regardless of anatomy_groups -- that argument only selects which surfaces + # are extracted. uniGradICON's Dice loss one-hots every class the two + # labelmaps share, on its 175^3 grid at batch 4, so handing it ~95 + # whole-body classes costs tens of gigabytes per step; the lung labels alone + # cost a twentieth of that and are the only ones the distance maps describe. + lung_label_ids = np.array( + sorted(segmenter.taxonomy.labels_in_group("lung")), dtype=np.uint16 + ) contour_tools = ContourTools(log_level=log_level) transform_tools = TransformTools() @@ -199,9 +206,13 @@ def segment_phase(image_file: Path) -> tuple[Path, Path]: contour_tools.save_combined_surfaces( segmentation_result["label_surfaces"], str(surface_file) ) - itk.imwrite( - segmentation_result["labelmap"], str(labelmap_file), compression=True + whole_body_labelmap = segmentation_result["labelmap"] + whole_body_arr = itk.GetArrayViewFromImage(whole_body_labelmap) + lung_labelmap = itk.GetImageFromArray( + np.where(np.isin(whole_body_arr, lung_label_ids), whole_body_arr, 0) ) + lung_labelmap.CopyInformation(whole_body_labelmap) + itk.imwrite(lung_labelmap, str(labelmap_file), compression=True) surface = cast(pv.PolyData, pv.read(str(surface_file))) # Rasterize the lung surface into ICON's distance-map representation, # mirroring ``RegisterModelsDistanceMaps._create_masks_from_models`` so @@ -437,8 +448,7 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: chain = RegisterImagesGreedyICON(log_level=log_level) chain.greedy.set_transform_type("Deformable") chain.greedy.set_metric("CC") - if number_of_iterations_greedy is not None: - chain.greedy.set_number_of_iterations(number_of_iterations_greedy) + chain.greedy.set_number_of_iterations(number_of_iterations_greedy) chain.icon.set_number_of_iterations(number_of_iterations_icon) chain.icon.set_mass_preservation(False) chain.icon.set_weights_path(str(method_weights)) @@ -448,8 +458,7 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: registrar.set_transform_type("Deformable") # CC is what RegisterModelsDistanceMaps uses on distance maps. registrar.set_metric("CC") - if number_of_iterations_greedy is not None: - registrar.set_number_of_iterations(number_of_iterations_greedy) + registrar.set_number_of_iterations(number_of_iterations_greedy) else: registrar = RegisterImagesICON(log_level=log_level) # None, not 0: icon_registration rejects 0 and takes None to mean diff --git a/tutorials/tutorial_02_lung_finetune_icon.py b/tutorials/tutorial_02_lung_finetune_icon.py index 89b13e8d..1e97b23c 100644 --- a/tutorials/tutorial_02_lung_finetune_icon.py +++ b/tutorials/tutorial_02_lung_finetune_icon.py @@ -4,7 +4,7 @@ Purpose ------- Finetune uniGradICON on every DIR-Lab 4D CT case except Case 1, then register -``Case1Pack_T00.mha`` (moving) to ``Case1Pack_T50.mha`` (fixed) four ways: +``Case1Pack_T00.mha`` (moving) to ``Case1Pack_T50.mha`` (fixed) several ways: ``RegisterImagesGreedy`` alone, deformable, with its default iteration schedule; ``RegisterImagesICON`` with the stock uniGradICON weights and with the finetuned weights; and ``RegisterImagesGreedyICON`` -- the same Greedy @@ -27,15 +27,46 @@ Reported per method: the mean, standard deviation, 95th percentile and maximum landmark error in millimeters; the mean, 5th percentile, median, 95th percentile, minimum and maximum of the per-class Dice scores; the number of -mislabeled voxels; and the wall-clock registration time. +mislabeled voxels; and the wall-clock registration time. The ``loss`` column is +*not* comparable across rows: each backend reports its own metric, and a chain +reports only its last stage's loss, measured against data the earlier stage +already warped. + +Why the chain does not win here +------------------------------- +On this pair ``RegisterImagesGreedyICON`` scores no better than Greedy alone, and +the extra rows exist to show why rather than to hide it. + +Compare a chain row against ``greedy_icon_stage0`` -- that row is the chain's own +Greedy stage scored on its own transform, and it is the right comparator because +Greedy is not bit-reproducible run to run (its scatter is a few thousandths of a +millimeter, the same size as the effect being measured). ``icon_residual_*`` is +how far the ICON stage moved the landmarks and ``tre_delta_*`` how much their +error changed as a result. + +What that shows: the ICON residual is small, and its per-landmark size is +essentially uncorrelated with how much it helps or hurts -- it is a random +perturbation of an already-better transform, not a correction. Raising ICON's +test-time optimization steps shrinks the residual and the damage together, so the +chain converges toward simply reproducing its Greedy stage, at several times the +runtime. The reason is resolution: ICON's residual deformation lives on +uniGradICON's fixed 175^3 network grid, which over this roughly 250mm field of +view is about 1.4mm between nodes -- coarser than the error Greedy has already +reached. A refinement stage that cannot resolve the error it is asked to remove +has nothing to contribute, and the chain applies its residual unconditionally. + +Ruled out as causes: the two stages are configured identically; the transform +composition was verified; every row is scored with the same landmarks in the same +direction; and the pre-warp between stages leaves under a percent of the fixed +grid without moving data, none of it near a landmark. Finetuning artifacts (dataset JSON, YAML config, checkpoint tree) are written under ``tutorials/network_weights/icon_dirlab_4dct``. The final checkpoint is ``tutorials/network_weights/icon_dirlab_4dct/icon_dirlab_4dct_model/checkpoints/ network_weights_final.trch``, the path returned by -``WorkflowFinetuneICONRegistration.expected_weights_path()``. That directory is -deleted at the start of every run, so each run finetunes from scratch; see the -comment above the ``shutil.rmtree`` call for how to reuse a previous run. +``WorkflowFinetuneICONRegistration.expected_weights_path()``. ``run_finetuning`` +is off by default so runs reuse that checkpoint; turning it on deletes the +directory and finetunes from scratch. Data Required ------------- @@ -92,6 +123,8 @@ weights_dir = tutorials_dir / "network_weights" finetune_name = "icon_dirlab_4dct" + # Set True to finetune from scratch. That deletes experiment_dir below, + # including any checkpoint a previous run left there. run_finetuning = True test_mode = TestTools.running_as_test() @@ -302,6 +335,8 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: moving_image, ReferenceImage=fixed_image, UseReferenceImage=True, + # Air, not ITK's default 0, which in CT is water. + DefaultPixelValue=-1000.0, ) unregistered_labelmap = itk.resample_image_filter( moving_labelmap, @@ -310,6 +345,33 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: UseReferenceImage=True, ) + # Diagnostic columns describing what the ICON stage of the chain added on + # top of its Greedy stage. They are empty on rows where they do not apply, + # but csv.DictWriter takes its field names from the first row, so every row + # has to carry the keys. + empty_chain_diagnostics: dict[str, Any] = { + "icon_residual_mean": None, + "icon_residual_p95": None, + "icon_residual_max": None, + "tre_delta_mean": None, + "tre_delta_max": None, + "tre_delta_residual_corr": None, + } + + def warp_moving(transform: itk.Transform) -> tuple[itk.Image, itk.Image]: + """Warp the moving image and labelmap onto the fixed grid.""" + return ( + transform_tools.transform_image( + moving_image, transform, fixed_image, background_value=-1000.0 + ), + transform_tools.transform_image( + moving_labelmap, + transform, + fixed_image, + interpolation_method="nearest", + ), + ) + registered_images: dict[str, itk.Image] = {"unregistered": unregistered_image} labelmaps: dict[str, itk.Image] = {"unregistered": unregistered_labelmap} rows: list[dict[str, Any]] = [ @@ -322,29 +384,42 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: np.linalg.norm(fixed_landmarks - moving_landmarks, axis=1) ), **overlap_metrics(unregistered_labelmap), + **empty_chain_diagnostics, } ] - for method_name, method_weights in ( - ("greedy", None), - ("icon_stock", None), - ("icon_finetuned", weights_path), - ("greedy_icon_finetuned", weights_path), - ): + + # Freezing Greedy and sweeping only the ICON stage's test-time optimization + # steps makes the residual the sole difference between the chain rows, which + # bounds how much test-time optimization could recover. + icon_step_sweep: list[Optional[int]] = [None, 10, 50] + methods: list[tuple[str, Optional[Path], Optional[int]]] = [ + ("greedy", None, None), + ("icon_stock", None, number_of_iterations_icon), + ("icon_finetuned", weights_path, number_of_iterations_icon), + ] + methods += [ + (f"greedy_icon_finetuned_steps_{steps}", weights_path, steps) + for steps in icon_step_sweep + ] + + greedy_stage_transform: Optional[itk.Transform] = None + for method_name, method_weights, icon_steps in methods: registrar: RegisterImagesBase + chain: Optional[RegisterImagesGreedyICON] = None if method_name == "greedy": registrar = RegisterImagesGreedy(log_level=log_level) registrar.set_transform_type("Deformable") if number_of_iterations_greedy is not None: registrar.set_number_of_iterations(number_of_iterations_greedy) - elif method_name == "greedy_icon_finetuned": + elif method_name.startswith("greedy_icon"): # Both stages are configured exactly as the standalone "greedy" and - # "icon_finetuned" rows above, so this row differs from + # "icon_finetuned" rows above, so these rows differ from # "icon_finetuned" only by the Greedy transform ICON starts from. chain = RegisterImagesGreedyICON(log_level=log_level) chain.greedy.set_transform_type("Deformable") if number_of_iterations_greedy is not None: chain.greedy.set_number_of_iterations(number_of_iterations_greedy) - chain.icon.set_number_of_iterations(number_of_iterations_icon) + chain.icon.set_number_of_iterations(icon_steps) chain.icon.set_mass_preservation(True) # For non-contrast CT chain.icon.set_weights_path(str(method_weights)) registrar = chain @@ -353,7 +428,7 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: # None, not 0: icon_registration rejects 0 and takes None to mean # "no test-time finetuning steps", so the comparison reflects what # each set of weights predicts rather than per-pair optimization. - registrar.set_number_of_iterations(number_of_iterations_icon) + registrar.set_number_of_iterations(icon_steps) registrar.set_mass_preservation(True) # For non-contrast CT if method_weights is not None: registrar.set_weights_path(str(method_weights)) @@ -364,26 +439,129 @@ def overlap_metrics(labelmap: itk.Image) -> dict[str, Any]: result = registrar.register(moving_image) elapsed_s = time.perf_counter() - start_time - registered_images[method_name] = transform_tools.transform_image( - moving_image, result["forward_transform"], fixed_image - ) - labelmaps[method_name] = transform_tools.transform_image( - moving_labelmap, - result["forward_transform"], - fixed_image, - interpolation_method="nearest", + composed_errors = landmark_errors(result["forward_transform"]) + chain_diagnostics = dict(empty_chain_diagnostics) + if chain is not None: + # RegisterImagesChain mirrors each stage's own result onto the + # sub-registrar before composing, so chain.greedy.forward_transform + # is the stage-only Greedy result and chain.icon.forward_transform + # is the residual ICON added on top of it. Both are exact + # transforms, scored the same way as every other row. + greedy_stage_errors = landmark_errors(chain.greedy.forward_transform) + icon_stage_transform: itk.Transform = chain.icon.forward_transform + residual_mm = np.array( + [ + np.linalg.norm( + np.asarray(icon_stage_transform.TransformPoint(tuple(point))) + - point + ) + for point in fixed_landmarks + ] + ) + delta_mm = composed_errors - greedy_stage_errors + chain_diagnostics = { + "icon_residual_mean": float(residual_mm.mean()), + "icon_residual_p95": float(np.percentile(residual_mm, 95)), + "icon_residual_max": float(residual_mm.max()), + "tre_delta_mean": float(delta_mm.mean()), + "tre_delta_max": float(delta_mm.max()), + "tre_delta_residual_corr": float( + np.corrcoef(delta_mm, residual_mm)[0, 1] + ), + } + + # The Greedy stage is identical across the sweep, so score it once. + if greedy_stage_transform is None: + greedy_stage_transform = chain.greedy.forward_transform + stage_image, stage_labelmap = warp_moving(greedy_stage_transform) + registered_images["greedy_icon_stage0"] = stage_image + labelmaps["greedy_icon_stage0"] = stage_labelmap + rows.append( + { + "method": "greedy_icon_stage0", + "weights": "-", + "registration_time_s": None, + "loss": None, + **landmark_metrics(greedy_stage_errors), + **overlap_metrics(stage_labelmap), + **empty_chain_diagnostics, + } + ) + + registered_images[method_name], labelmaps[method_name] = warp_moving( + result["forward_transform"] ) rows.append( { "method": method_name, "weights": str(method_weights) if method_weights else "-", "registration_time_s": elapsed_s, + # Not comparable across rows: each backend reports its own + # metric, and a chain reports only its last stage's loss, + # measured against data the earlier stage already warped. "loss": float(result["loss"]), - **landmark_metrics(landmark_errors(result["forward_transform"])), + **landmark_metrics(composed_errors), **overlap_metrics(labelmaps[method_name]), + **chain_diagnostics, } ) + # How much of the fixed grid the pre-warp has no moving data for. Those + # voxels used to be filled with ITK's default 0 -- water in CT -- which the + # ICON stage then saw as tissue where the moving image had nothing. + if greedy_stage_transform is not None: + coverage_input = itk.image_from_array( + np.ones(list(itk.size(moving_image))[::-1], dtype=np.float32) + ) + coverage_input.CopyInformation(moving_image) + coverage = transform_tools.transform_image( + coverage_input, greedy_stage_transform, fixed_image, background_value=0.0 + ) + itk.imwrite( + coverage, str(output_dir / "prewarp_coverage.mha"), compression=True + ) + coverage_arr = itk.GetArrayFromImage(coverage) + uncovered = coverage_arr < 0.999 + fixed_hu = itk.GetArrayFromImage(fixed_image) + uncovered_in_air = uncovered & (fixed_hu < -500.0) + reporter.log_info( + "Pre-warp coverage: %d/%d fixed voxels uncovered (%.4f%%), " + "%d of them where the fixed image is air", + int(uncovered.sum()), + uncovered.size, + 100.0 * uncovered.sum() / uncovered.size, + int(uncovered_in_air.sum()), + ) + if uncovered.any(): + k_counts = uncovered.sum(axis=(1, 2)) + reporter.log_info( + "Uncovered voxels per slice along k: min %d, max %d, " + "first slice %d, last slice %d", + int(k_counts.min()), + int(k_counts.max()), + int(k_counts[0]), + int(k_counts[-1]), + ) + # Landmarks within 10mm of missing data are the ones a bad fill + # value could plausibly have moved. + spacing = np.asarray(fixed_image.GetSpacing(), dtype=np.float64) + radius = np.maximum(1, np.ceil(10.0 / spacing)).astype(int) + near_count = 0 + for point in fixed_landmarks: + index = fixed_image.TransformPhysicalPointToIndex(tuple(point)) + lo = [max(0, int(index[d]) - int(radius[d])) for d in range(3)] + hi = [ + min(uncovered.shape[2 - d], int(index[d]) + int(radius[d]) + 1) + for d in range(3) + ] + if uncovered[lo[2] : hi[2], lo[1] : hi[1], lo[0] : hi[0]].any(): + near_count += 1 + reporter.log_info( + "Landmarks within 10mm of uncovered data: %d/%d", + near_count, + len(fixed_landmarks), + ) + # Result saving itk.imwrite( fixed_labelmap, str(output_dir / "fixed_labelmap.mha"), compression=True diff --git a/tutorials/tutorial_03_heart_reconstruct_highres_4d_ct.py b/tutorials/tutorial_03_heart_reconstruct_highres_4d_ct.py index c39bcc2d..249eef14 100644 --- a/tutorials/tutorial_03_heart_reconstruct_highres_4d_ct.py +++ b/tutorials/tutorial_03_heart_reconstruct_highres_4d_ct.py @@ -23,7 +23,7 @@ import itk from physiotwin4d import ( - RegisterImagesGreedyICON, + RegisterImagesGreedy, TestTools, WorkflowReconstructHighres4DCT, ) @@ -57,9 +57,8 @@ log_level = logging.INFO - registration_method = RegisterImagesGreedyICON(log_level=log_level) - registration_method.greedy.set_number_of_iterations(number_of_iterations_greedy) - registration_method.icon.set_mass_preservation(False) # For contrast-enhanced CT + registration_method = RegisterImagesGreedy(log_level=log_level) + registration_method.set_number_of_iterations(number_of_iterations_greedy) # Directory setup and data reading diff --git a/tutorials/tutorial_03_lung_reconstruct_highres_4d_ct.py b/tutorials/tutorial_03_lung_reconstruct_highres_4d_ct.py index 73378dd1..137c0c7a 100644 --- a/tutorials/tutorial_03_lung_reconstruct_highres_4d_ct.py +++ b/tutorials/tutorial_03_lung_reconstruct_highres_4d_ct.py @@ -12,9 +12,6 @@ ------------- Full data: ``data/DirLab-4DCT/Case1Pack_T??.mha`` Test data: ``data/test/DirLab-4DCT/Case1Pack_T??.mha`` -ICON weights: Tutorial 2 output -(``tutorials/network_weights/icon_dirlab_4dct/.../network_weights_final.trch``), -optional — the stock uniGradICON weights are used when it is absent. Outputs (under ``tutorials/output/tutorial_03_lung/``) ----------------------------------------------------- @@ -32,7 +29,7 @@ import itk from physiotwin4d import ( - RegisterImagesGreedyICON, + RegisterImagesGreedy, TestTools, WorkflowReconstructHighres4DCT, ) @@ -58,17 +55,6 @@ # data/DirLab-4DCT/fix_downloaded_data.py. case_glob = "Case1Pack_T??.mha" - # Weights finetuned on DIR-Lab by Tutorial 2; see - # WorkflowFinetuneICONRegistration.expected_weights_path(). - icon_weights_path = ( - tutorials_dir - / "network_weights" - / "icon_dirlab_4dct" - / "icon_dirlab_4dct_model" - / "checkpoints" - / "network_weights_final.trch" - ) - test_mode = TestTools.running_as_test() if test_mode: data_dir = repo_root / "data" / "test" / "DirLab-4DCT" @@ -79,29 +65,12 @@ log_level = logging.INFO - logging.basicConfig(level=log_level) - logger = logging.getLogger(class_name) - # Directory setup and data reading output_dir.mkdir(parents=True, exist_ok=True) - # Registration: Greedy for the coarse motion, then ICON for the deformable - # refinement. The Tutorial 2 weights are used when they exist; without them - # the tutorial still runs, on the stock uniGradICON weights. - registration_method = RegisterImagesGreedyICON(log_level=log_level) - registration_method.greedy.set_number_of_iterations(number_of_iterations_greedy) - registration_method.icon.set_mass_preservation(True) # For non-contrast CT - if icon_weights_path.exists(): - registration_method.icon.set_weights_path(str(icon_weights_path)) - logger.info("Registering with finetuned ICON weights: %s", icon_weights_path) - else: - logger.warning( - "Finetuned ICON weights not found at %s; registering with the stock " - "uniGradICON weights. Run " - "tutorials/tutorial_02_lung_finetune_icon.py to create them.", - icon_weights_path, - ) + registration_method = RegisterImagesGreedy(log_level=log_level) + registration_method.set_number_of_iterations(number_of_iterations_greedy) phase_files = sorted(data_dir.glob(case_glob)) if not phase_files: diff --git a/tutorials/tutorial_06_heart_create_statistical_model.py b/tutorials/tutorial_06_heart_create_statistical_model.py index b62a8499..864f4c40 100644 --- a/tutorials/tutorial_06_heart_create_statistical_model.py +++ b/tutorials/tutorial_06_heart_create_statistical_model.py @@ -22,6 +22,7 @@ import numpy as np import pyvista as pv +from parameters_heart_ct_kcl import HEART_CT_KCL from physiotwin4d import ( TestTools, @@ -48,10 +49,9 @@ test_mode = TestTools.running_as_test() if test_mode: data_dir = repo_root / "data" / "test" / "KCL-Heart-Model" - number_of_pca_components = 5 else: data_dir = repo_root / "data" / "KCL-Heart-Model" - number_of_pca_components = 10 + number_of_pca_components = HEART_CT_KCL.pca_components(test_mode) log_level = logging.INFO diff --git a/tutorials/tutorial_06_lung_create_statistical_model.py b/tutorials/tutorial_06_lung_create_statistical_model.py index 8c475e6b..c2d8ec45 100644 --- a/tutorials/tutorial_06_lung_create_statistical_model.py +++ b/tutorials/tutorial_06_lung_create_statistical_model.py @@ -39,6 +39,7 @@ import itk import numpy as np import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( ContourTools, @@ -68,7 +69,9 @@ data_dir = repo_root / "data" / "DirLab-4DCT" - number_of_pca_components = 5 + number_of_pca_components = LUNG_CT_DIRLAB.pca_components( + TestTools.running_as_test() + ) # Atlas iterations used to build the reference surface; 1 is a single # template-biased pass. diff --git a/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py index 36d54905..8e446a92 100644 --- a/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py +++ b/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py @@ -23,6 +23,7 @@ import itk import pyvista as pv +from parameters_heart_ct_kcl import HEART_CT_KCL from physiotwin4d import ( ContourTools, @@ -63,6 +64,20 @@ data_dir = repo_root / "data" / "DirLab-4DCT" patient_image_file = data_dir / "Case1Pack_T70.mha" + # Distance-map weights finetuned by + # tutorial_02_duke_heart_distancemap_finetune_icon.py. The heart has its own + # finetuning run rather than reusing the lung one's: the heart registration + # mask is far tighter, so heart distance maps saturate over a shorter radius + # and do not share an intensity distribution with lung ones. + icon_weights_path = ( + tutorials_dir + / "network_weights" + / "icon_duke_heart_distancemap" + / "icon_duke_heart_distancemap_model" + / "checkpoints" + / "network_weights_final.trch" + ) + log_level = logging.INFO segmentation_method = SegmentChestTotalSegmentator() @@ -140,9 +155,12 @@ patient_image=patient_image, patient_labelmap=heart_labelmap, log_level=log_level, - labelmap_interior_object_ids=[141, 142, 143, 144], - # These are the internal chambers of the heart when using TotalSegmentator. + # This patient is segmented with TotalSegmentator, so its chamber ids + # are the ones to keep out of the distance map. + labelmap_interior_object_ids=HEART_CT_KCL.interior_object_ids_totalsegmentator, ) + workflow.set_mask_dilation_mm(HEART_CT_KCL.mask_dilation_mm) + workflow.set_distancemap_squared_max(HEART_CT_KCL.distancemap_squared_max) if pca_model is not None: workflow.set_use_pca_registration( use_pca_registration=True, @@ -150,6 +168,20 @@ use_surface=False, ) + # The labelmap-to-labelmap stage registers distance maps, not intensities, + # so it uses the distance-map-finetuned weights when they exist; without + # them the tutorial still runs, on the stock uniGradICON weights. + if icon_weights_path.exists(): + workflow.set_labelmap_to_labelmap_icon_weights_path(str(icon_weights_path)) + else: + workflow.log_warning( + "Finetuned distance-map ICON weights not found at %s; fitting with " + "the stock uniGradICON weights. Run " + "tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py to " + "create them.", + icon_weights_path, + ) + # Workflow execution workflow_results = workflow.process() diff --git a/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py index 20614dd5..5a33a947 100644 --- a/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py +++ b/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py @@ -30,6 +30,7 @@ import itk import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( ContourTools, @@ -61,14 +62,15 @@ pca_json = tutorial_06_dir / "pca_model.json" pca_mean_file = tutorial_06_dir / "pca_mean_surface.vtp" - number_of_pca_components = 5 + number_of_pca_components = LUNG_CT_DIRLAB.pca_components( + TestTools.running_as_test() + ) patient_image_file = repo_root / "data" / "Chest-CT" / "Chest-CT.mha" - # Distance-map weights finetuned on DIR-Lab by Tutorial 2; see - # WorkflowFinetuneICONRegistration.expected_weights_path(). The - # mask_dilation_mm set below must match the one that tutorial finetuned - # with, since it fixes the distance maps' saturation radius. + # Distance-map weights finetuned on DIR-Lab by + # tutorial_02_lung_distancemap_finetune_icon.py; see + # WorkflowFinetuneICONRegistration.expected_weights_path(). icon_weights_path = ( tutorials_dir / "network_weights" @@ -145,7 +147,8 @@ use_surface=False, ) - workflow.set_mask_dilation_mm(mask_dilation_mm=40) + workflow.set_mask_dilation_mm(LUNG_CT_DIRLAB.mask_dilation_mm) + workflow.set_distancemap_squared_max(LUNG_CT_DIRLAB.distancemap_squared_max) # The labelmap-to-labelmap stage registers distance maps, not intensities, # so it uses the distance-map-finetuned weights when they exist; without diff --git a/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py b/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py index f40ec176..77af9dad 100644 --- a/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py +++ b/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py @@ -14,9 +14,8 @@ the fitted SSM surface, sharing the model's fixed topology. 2. Propagate the fitted surface to every respiratory phase. Each phase is - registered to the reference phase with ``RegisterImagesGreedyICON`` - (``WorkflowReconstructHighres4DCT``), using the DIR-Lab weights finetuned by - ``tutorial_02_lung_finetune_icon.py``. The forward transform for each phase + registered to the reference phase with ``RegisterImagesGreedy`` + (``WorkflowReconstructHighres4DCT``). The forward transform for each phase warps the fitted SSM surface (``TransformTools.transform_pvcontour``, with deformation magnitude attached), producing one ``*_T{PP}_ssm_surface.vtp`` per phase. @@ -31,12 +30,10 @@ PCA model: Tutorial 6 output (``output/tutorial_06_lung/pca_model.json``, ``pca_mean_surface.vtp``) ICON weights: Tutorial 2 output -(``network_weights/icon_dirlab_4dct/icon_dirlab_4dct_model/checkpoints/ -network_weights_final.trch``) for the phase-to-phase CT registration, and -``network_weights/icon_dirlab_4dct_distancemap/ -icon_dirlab_4dct_distancemap_model/checkpoints/network_weights_final.trch`` for -the distance-map stage of the SSM fit. Both optional — the stock uniGradICON -weights are used when they are absent. +(``network_weights/icon_dirlab_4dct_distancemap/ +icon_dirlab_4dct_distancemap_model/checkpoints/network_weights_final.trch``) for +the distance-map stage of the SSM fit. Optional — the stock uniGradICON weights +are used when it is absent. Outputs (per case, under ``output/tutorial_08_lung//``) ------------------------------------------------------------ @@ -57,10 +54,11 @@ import itk import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( ContourTools, - RegisterImagesGreedyICON, + RegisterImagesGreedy, SegmentNVSegmentCTMRI, TestTools, TransformTools, @@ -93,22 +91,9 @@ pca_model_file = tutorial_06_dir / "pca_model.json" pca_mean_file = tutorial_06_dir / "pca_mean_surface.vtp" - # Weights finetuned on DIR-Lab by Tutorial 2; see - # WorkflowFinetuneICONRegistration.expected_weights_path(). - icon_weights_path = ( - tutorials_dir - / "network_weights" - / "icon_dirlab_4dct" - / "icon_dirlab_4dct_model" - / "checkpoints" - / "network_weights_final.trch" - ) - # Distance-map weights finetuned on DIR-Lab by # tutorial_02_lung_distancemap_finetune_icon.py, used by the - # labelmap-to-labelmap stage of the SSM fit. mask_dilation_mm below must - # match the value that tutorial finetuned with, since it fixes the distance - # maps' saturation radius. + # labelmap-to-labelmap stage of the SSM fit. icon_distancemap_weights_path = ( tutorials_dir / "network_weights" @@ -117,7 +102,10 @@ / "checkpoints" / "network_weights_final.trch" ) - fit_mask_dilation_mm = 40.0 + + number_of_pca_components = LUNG_CT_DIRLAB.pca_components( + TestTools.running_as_test() + ) # Phase the SSM is fitted to; Tutorial 6 builds the lung PCA model from the # T70 surfaces, so the same phase is used here as the fitting reference. @@ -142,19 +130,8 @@ with pca_model_file.open(encoding="utf-8") as f: pca_model = json.load(f) - # The Tutorial 2 weights are used when they exist; without them the tutorial - # still runs, on the stock uniGradICON weights. - use_finetuned_icon_weights = icon_weights_path.exists() - if use_finetuned_icon_weights: - logger.info("Registering with finetuned ICON weights: %s", icon_weights_path) - else: - logger.warning( - "Finetuned ICON weights not found at %s; registering with the stock " - "uniGradICON weights. Run " - "tutorials/tutorial_02_lung_finetune_icon.py to create them.", - icon_weights_path, - ) - + # The Tutorial 2 distance-map weights are used when they exist; without them + # the tutorial still runs, on the stock uniGradICON weights. use_finetuned_distancemap_weights = icon_distancemap_weights_path.exists() if use_finetuned_distancemap_weights: logger.info( @@ -233,10 +210,11 @@ fit_workflow.set_use_pca_registration( use_pca_registration=True, pca_model=pca_model, - number_of_pca_components=6, + number_of_pca_components=number_of_pca_components, use_surface=False, ) - fit_workflow.set_mask_dilation_mm(mask_dilation_mm=fit_mask_dilation_mm) + fit_workflow.set_mask_dilation_mm(LUNG_CT_DIRLAB.mask_dilation_mm) + fit_workflow.set_distancemap_squared_max(LUNG_CT_DIRLAB.distancemap_squared_max) if use_finetuned_distancemap_weights: fit_workflow.set_labelmap_to_labelmap_icon_weights_path( str(icon_distancemap_weights_path) @@ -269,10 +247,7 @@ phase_ids = [path.stem.split("_")[1] for path in phase_files] time_series = [itk.imread(str(path)) for path in phase_files] - registration_method = RegisterImagesGreedyICON(log_level=log_level) - if use_finetuned_icon_weights: - registration_method.icon.set_weights_path(str(icon_weights_path)) - registration_method.icon.set_mass_preservation(True) # For non-contrast CT + registration_method = RegisterImagesGreedy(log_level=log_level) reg_workflow = WorkflowReconstructHighres4DCT( time_series_images=time_series, From 18d225e7d0449eea50af4110e6e36c17bc937b2e Mon Sep 17 00:00:00 2001 From: Stephen Aylward Date: Tue, 11 Aug 2026 14:45:03 -0400 Subject: [PATCH 2/3] ENH: Label-aware meshing, labeled USD export, Duke heart tutorial track ContourTools gains the mesh generation the labelmap tutorials need: extract_watertight_surface (pad + blur + flying edges, outward normals), extract_label_surfaces (per-label surfaces that conform at shared walls via signed distance maps), extract_tetrahedra + trim_tetrahedra_to_surface, remesh_and_smooth_surface (ACVD clustering + Taubin), and apply_anatomy_color. extract_contours now resamples anisotropic labelmaps onto an isotropic grid and Taubin-smooths, so contours no longer terrace on the voxel pitch. Adds pyacvd dependency. USD export: WorkflowConvertVTKToUSD takes label_names/segmenter and splits each mesh on its per-cell label array, giving one prim per structure at /World/{project}/{group}/{structure} so identity survives a time series. Labels now take precedence over static_merge (combining both raises). UsdMeshConverter time-samples faceVertexCounts/Indices when frames disagree on topology, instead of silently reusing frame 0's triangulation. USDTools.get_mesh_paths descends the full subtree. Registration: out-of-FOV voxels fill with the modality background value rather than 0 in get_registered_image and the 4D reconstruction; transform_tools validates background_value against the pixel type range and handles bool images. ImageTools.resample_image_by_scale resamples by voxel-count multiplier while preserving physical extent. Breaking: surface_target_reduction renamed to surface_reduction_rate (WorkflowConvertImageToVTK.process, convert-image-to-vtk CLI flag); WorkflowConvertImageToUSD gains surface_reduction_rate. Tutorials: new Duke heart labelmap track (parameters module plus tutorials 04-10), per-organ parameter modules updated, docs/tutorials.rst and tutorials/README.md refreshed, LFS + gitignore rules for the lung MGN weights. Tests: test_contour_mesh_extraction.py, test_register_images_base.py, plus coverage for labeled and topology-varying USD export. --- .gitattributes | 4 + .gitignore | 3 +- docs/assets/tutorial_04_duke_heart.png | 3 + docs/assets/tutorial_04_heart-2png.png | 3 + docs/assets/tutorial_04_lung-2.png | 3 + docs/tutorials.rst | 24 +- pyproject.toml | 8 + src/physiotwin4d/cli/convert_image_to_vtk.py | 4 +- src/physiotwin4d/contour_tools.py | 749 +++++++++++++++++- src/physiotwin4d/convert_vtk_to_usd.py | 22 +- src/physiotwin4d/image_tools.py | 59 ++ src/physiotwin4d/register_images_base.py | 9 +- .../register_time_series_images.py | 9 +- src/physiotwin4d/transform_tools.py | 22 +- src/physiotwin4d/usd_tools.py | 16 +- .../vtk_to_usd/usd_mesh_converter.py | 93 ++- .../workflow_convert_image_to_usd.py | 17 +- .../workflow_convert_image_to_vtk.py | 16 +- .../workflow_convert_vtk_to_usd.py | 134 +++- ...rkflow_fit_statistical_model_to_patient.py | 3 + .../workflow_reconstruct_highres_4d_ct.py | 8 +- tests/test_contour_mesh_extraction.py | 401 ++++++++++ tests/test_convert_vtk_to_usd.py | 63 +- tests/test_image_tools.py | 54 ++ tests/test_register_images_base.py | 116 +++ tests/test_tutorials.py | 13 +- tests/test_workflow_convert_vtk_to_usd.py | 151 +++- tutorials/README.md | 7 + tutorials/parameters_duke_heart_labelmaps.py | 152 ++++ tutorials/parameters_heart_ct_kcl.py | 61 +- tutorials/parameters_lung_ct_dirlab.py | 66 +- .../tutorial_01_heart_gated_ct_to_usd.py | 2 + tutorials/tutorial_01_lung_gated_ct_to_usd.py | 2 + ...02_duke_heart_distancemap_finetune_icon.py | 61 +- .../tutorial_04_duke_heart_labelmap_to_vtk.py | 356 +++++++++ tutorials/tutorial_04_heart_ct_to_vtk.py | 5 +- tutorials/tutorial_04_lung_ct_to_vtk.py | 6 +- .../tutorial_05_duke_heart_vtk_to_usd.py | 155 ++++ ..._06_duke_heart_create_statistical_model.py | 231 ++++++ ...orial_06_heart_create_statistical_model.py | 19 +- ...torial_06_lung_create_statistical_model.py | 24 +- ..._heart_fit_statistical_model_to_patient.py | 250 ++++++ ..._heart_fit_statistical_model_to_patient.py | 20 +- ...7_lung_fit_statistical_model_to_patient.py | 16 +- ..._08_duke_heart_fit_model_to_4d_patients.py | 362 +++++++++ ...torial_08_lung_fit_model_to_4d_patients.py | 8 +- ...ial_09_duke_heart_train_physicsnemo_mgn.py | 337 ++++++++ .../tutorial_09_lung_train_physicsnemo_mgn.py | 48 +- ...utorial_10_duke_heart_infer_physicsnemo.py | 202 +++++ .../tutorial_10_lung_infer_physicsnemo_mgn.py | 35 +- 50 files changed, 4214 insertions(+), 218 deletions(-) create mode 100644 docs/assets/tutorial_04_duke_heart.png create mode 100644 docs/assets/tutorial_04_heart-2png.png create mode 100644 docs/assets/tutorial_04_lung-2.png create mode 100644 tests/test_contour_mesh_extraction.py create mode 100644 tests/test_register_images_base.py create mode 100644 tutorials/parameters_duke_heart_labelmaps.py create mode 100644 tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py create mode 100644 tutorials/tutorial_05_duke_heart_vtk_to_usd.py create mode 100644 tutorials/tutorial_06_duke_heart_create_statistical_model.py create mode 100644 tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py create mode 100644 tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py create mode 100644 tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py create mode 100644 tutorials/tutorial_10_duke_heart_infer_physicsnemo.py diff --git a/.gitattributes b/.gitattributes index 0bbbfa0c..fd851d0f 100644 --- a/.gitattributes +++ b/.gitattributes @@ -6,3 +6,7 @@ tutorials/network_weights/physicsnemo_mgn_heart/*.pt filter=lfs diff=lfs merge=l tutorials/network_weights/physicsnemo_mgn_heart/*.vtu filter=lfs diff=lfs merge=lfs -text tutorials/network_weights/physicsnemo_mgn_heart/*.vtp filter=lfs diff=lfs merge=lfs -text tutorials/network_weights/physicsnemo_mgn_heart/pca_model.json filter=lfs diff=lfs merge=lfs -text +tutorials/network_weights/physicsnemo_mgn_lung_motion/*.pt filter=lfs diff=lfs merge=lfs -text +tutorials/network_weights/physicsnemo_mgn_lung_motion/*.vtu filter=lfs diff=lfs merge=lfs -text +tutorials/network_weights/physicsnemo_mgn_lung_motion/*.vtp filter=lfs diff=lfs merge=lfs -text +tutorials/network_weights/physicsnemo_mgn_lung_motion/pca_model.json filter=lfs diff=lfs merge=lfs -text diff --git a/.gitignore b/.gitignore index aae1e062..6627d71f 100644 --- a/.gitignore +++ b/.gitignore @@ -34,10 +34,11 @@ docs/assets/*.mp4 # Network weights network_weights -# Track the MGN heart model weights (other network_weights stay ignored) +# Track the MGN heart and lung model weights (other network_weights stay ignored) !tutorials/network_weights/ tutorials/network_weights/* !tutorials/network_weights/physicsnemo_mgn_heart/ +!tutorials/network_weights/physicsnemo_mgn_lung_motion/ # Data files *.gz diff --git a/docs/assets/tutorial_04_duke_heart.png b/docs/assets/tutorial_04_duke_heart.png new file mode 100644 index 00000000..3a7dbc93 --- /dev/null +++ b/docs/assets/tutorial_04_duke_heart.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:12f88e11e5f33668fb6fab896327e90111dbbd0a792248dff412eabfb45024d0 +size 953287 diff --git a/docs/assets/tutorial_04_heart-2png.png b/docs/assets/tutorial_04_heart-2png.png new file mode 100644 index 00000000..99dedbf1 --- /dev/null +++ b/docs/assets/tutorial_04_heart-2png.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc39cf799b47fb67b9ae886a24695da6a67742cd064daa88ab03358efe38ada6 +size 1397942 diff --git a/docs/assets/tutorial_04_lung-2.png b/docs/assets/tutorial_04_lung-2.png new file mode 100644 index 00000000..2b39d5ab --- /dev/null +++ b/docs/assets/tutorial_04_lung-2.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69899596e108eb82535ffeae3c24809e09d089da6e903ba2d5871397874d714d +size 804914 diff --git a/docs/tutorials.rst b/docs/tutorials.rst index d505f601..03e1d807 100644 --- a/docs/tutorials.rst +++ b/docs/tutorials.rst @@ -132,13 +132,15 @@ pipeline on top. 1. **Tutorial 1** — after downloading Slicer-Heart-CT. 2. **Tutorial 2** — after obtaining DIR-Lab. It writes the finetuned ICON - weights that Tutorials 3 (lung) and 8 use. -3. **Tutorial 3** — after Tutorial 2, whose weights it registers with. + weights Tutorial 8 uses when present; optional if stock weights are + acceptable. +3. **Tutorial 3** — after obtaining its dataset; it registers with Greedy and + needs no finetuned weights. 4. **Tutorial 4** — after downloading Slicer-Heart-CT. 5. **Tutorial 5** — after Tutorial 4, whose surfaces it converts. 6. **Tutorial 6** — heart needs KCL-Heart-Model, lung needs DIR-Lab. 7. **Tutorial 7** — after Tutorial 6; the lung variant also needs Chest-CT. -8. **Tutorial 8** — after Tutorials 2 and 6 (lung). +8. **Tutorial 8** — after Tutorial 6 (lung); Tutorial 2 optional. 9. **Tutorial 9** — after Tutorial 8, whose fitted meshes it trains on. 10. **Tutorial 10** — after Tutorial 9, whose checkpoint it loads. @@ -415,7 +417,7 @@ Inner API usage ) result = workflow.process( input_image=ct_image, - surface_target_reduction=0.5, + surface_reduction_rate=HEART_CT_KCL.surface_reduction_rate, extract_label_surfaces=save_label_surfaces, ) @@ -435,7 +437,8 @@ Adapt to your data Change the input volume path, then choose the segmenter matching your scan: contrast versus non-contrast CT, or :class:`~physiotwin4d.SegmentNVSegmentCTMRI` for CT **and** MRI. Raise - ``surface_target_reduction`` toward ``1.0`` for lighter meshes. Every + ``surface_reduction_rate`` in the tutorial's parameter module toward ``1.0`` + for lighter meshes. Every segmenter declares its own labels through :class:`~physiotwin4d.AnatomyTaxonomy`, so downstream grouping and USD materials follow automatically — see :doc:`api/segmentation/index`. @@ -764,10 +767,11 @@ Run python tutorials/tutorial_09_lung_train_physicsnemo_mgn.py Outputs - ``mgn_stage_model.pt``, its metadata and loss/RMSE logs, the per-case - manifests, and the held-out evaluation under ``eval_mgn/`` — in the - directory training used: ``tutorials/output/tutorial_09_lung_mgn/`` - normally, or a fresh sibling when resuming. + ``mgn_stage_model.pt``, its metadata and loss/RMSE logs, in the shared + weights directory Tutorial 10 reads + (``tutorials/network_weights/physicsnemo_mgn_lung_motion/``, a fresh sibling + of it when resuming). The per-case manifests and the held-out evaluation + under ``eval_mgn/`` stay in ``tutorials/output/tutorial_09_lung_mgn/``. Adapt to your data The contract is the manifest, not the tutorial. Each JSON names a reference @@ -827,7 +831,7 @@ Run Outputs The predicted surface, its error statistics against the ground-truth phase in millimetres, and a USD scene, under - ``tutorials/output/tutorial_09_lung_mgn/tutorial_10_lung_mgn//``. + ``tutorials/output/tutorial_10_lung_mgn//``. Adapt to your data Change ``case_id`` and ``stage_fraction`` to predict a different subject, or diff --git a/pyproject.toml b/pyproject.toml index aa20f911..b75d518c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -85,6 +85,7 @@ dependencies = [ "pyvista[all]>=0.47.0", "usd-core>=23.11", "trimesh>=4.0.0", + "pyacvd>=0.4.0", # Utilities "ipykernel>=6.0.0", @@ -289,6 +290,8 @@ module = [ "physicsnemo.*", "pxr", "pxr.*", + "pyacvd", + "pyacvd.*", "pyvista", "pyvista.*", "simpleware", @@ -325,6 +328,7 @@ module = [ # pre-commit mypy environment and ships no py.typed marker when installed. # mypy per-module patterns match whole components only, so each tutorial is # listed by name; add a line here when adding or renaming a tutorial. + "parameters_duke_heart_labelmaps", "parameters_heart_ct_kcl", "parameters_lung_ct_dirlab", "tutorial_01_heart_gated_ct_to_usd", @@ -334,11 +338,15 @@ module = [ "tutorial_02_lung_finetune_icon", "tutorial_03_heart_reconstruct_highres_4d_ct", "tutorial_03_lung_reconstruct_highres_4d_ct", + "tutorial_04_duke_heart_labelmap_to_vtk", "tutorial_04_heart_ct_to_vtk", "tutorial_04_lung_ct_to_vtk", + "tutorial_05_duke_heart_vtk_to_usd", "tutorial_05_heart_vtk_to_usd", + "tutorial_06_duke_heart_create_statistical_model", "tutorial_06_heart_create_statistical_model", "tutorial_06_lung_create_statistical_model", + "tutorial_07_duke_heart_fit_statistical_model_to_patient", "tutorial_07_heart_fit_statistical_model_to_patient", "tutorial_07_lung_fit_statistical_model_to_patient", "tutorial_08_lung_fit_model_to_4d_patients", diff --git a/src/physiotwin4d/cli/convert_image_to_vtk.py b/src/physiotwin4d/cli/convert_image_to_vtk.py index 43032469..245af88e 100644 --- a/src/physiotwin4d/cli/convert_image_to_vtk.py +++ b/src/physiotwin4d/cli/convert_image_to_vtk.py @@ -110,7 +110,7 @@ def main() -> int: ), ) parser.add_argument( - "--surface-target-reduction", + "--surface-reduction-rate", type=float, default=0.0, help=( @@ -177,7 +177,7 @@ def main() -> int: result = workflow.process( input_image=input_image, anatomy_groups=args.anatomy_groups, - surface_target_reduction=args.surface_target_reduction, + surface_reduction_rate=args.surface_reduction_rate, extract_label_surfaces=(args.output_mode == "label"), ) except (ValueError, RuntimeError, OSError) as exc: diff --git a/src/physiotwin4d/contour_tools.py b/src/physiotwin4d/contour_tools.py index 8ec8c0c1..c0539301 100644 --- a/src/physiotwin4d/contour_tools.py +++ b/src/physiotwin4d/contour_tools.py @@ -6,16 +6,41 @@ import logging import os -from typing import cast +from typing import Optional, Sequence, cast import itk import numpy as np +import pyacvd import pyvista as pv import trimesh from .image_tools import ImageTools from .physiotwin4d_base import PhysioTwin4DBase from .transform_tools import TransformTools +from .usd_anatomy_tools import USDAnatomyTools + +# VTK_VOXEL lists its eight corners in (i, j, k) raster order; VTK_HEXAHEDRON +# wants the bottom quad wound consistently, then the matching top quad. +_VOXEL_TO_HEX = [0, 1, 3, 2, 4, 5, 7, 6] +# When the image direction matrix is left-handed (negative determinant), the +# above winding yields inverted (negative-Jacobian) hexahedra, so swap the +# bottom and top quads to restore a positive volume. +_VOXEL_TO_HEX_FLIPPED = [4, 5, 7, 6, 0, 1, 3, 2] + +# trim_tetrahedra_to_surface halves a point's move until its cells clear the +# quality bound; below this fraction of the original move the point is put back +# where it started instead, which both ends the search and keeps the mesh free +# of moves too small to matter. +_MIN_TRIM_DAMPING = 2.0**-10 + +# Background voxels kept around the labels in extract_label_surfaces, enough +# for the isotropic grid to hold a full voxel of background on every side. +_LABEL_SURFACE_PAD = 3 + +# extract_contours resamples a labelmap whose coarsest spacing is more than +# this multiple of its finest one, because a contour built on the coarse axis +# terraces at its pitch. Below it the resample costs more than it buys. +_CONTOUR_ANISOTROPY_LIMIT = 1.5 class ContourTools(PhysioTwin4DBase): @@ -31,21 +56,106 @@ def __init__(self, log_level: int | str = logging.INFO): """ super().__init__(class_name=self.__class__.__name__, log_level=log_level) - @staticmethod + # USDAnatomyTools builds its color tables in __init__ without touching + # the stage, so stage=None is safe for these lookup-only uses. + self._anatomy_tools = USDAnatomyTools(stage=None, log_level=log_level) + + def apply_anatomy_color( + self, mesh: pv.DataSet, anatomy_names: Sequence[str] + ) -> None: + """Attach a structure's :class:`USDAnatomyTools` color **in-place**. + + Sets, as :meth:`WorkflowConvertImageToVTK._annotate` does, so geometry + from here colors the same way in Paraview, PyVista, and the USD + exporter: + + - ``field_data['AnatomyColor']`` — RGB float32 color. + - ``cell_data['Color']`` — RGBA uint8 solid color (n_cells × 4). + + Args: + mesh: Surface or volume mesh to annotate. + anatomy_names: Names tried in order, most specific first, e.g. an + organ name followed by its anatomy group. ``USDAnatomyTools`` + carries overrides for some organs (``myocardium``) but not + others (``left_ventricle``), so a group name is the usual + second entry. Falls back to ``'other'`` when none resolves. + """ + name = next( + ( + candidate + for candidate in anatomy_names + if self._anatomy_tools.resolve_anatomy_type(candidate) is not None + ), + None, + ) + if name is None: + self.log_warning( + "No anatomy color matches %s; using 'other'", + " or ".join(anatomy_names), + ) + name = "other" + + color_rgb = self._anatomy_tools.get_anatomy_diffuse_color(name) + mesh.field_data["AnatomyColor"] = np.array(color_rgb, dtype=np.float32) + rgba = np.array( + [int(channel * 255) for channel in color_rgb] + [255], dtype=np.uint8 + ) + if mesh.n_cells > 0: + mesh.cell_data["Color"] = np.tile(rgba, (mesh.n_cells, 1)) + def extract_contours( + self, labelmap_image: itk.image, smoothing_iterations: int = 10, smoothing_scale: float = 1.0, + surface_reduction_rate: float = 0.0, + taubin_iterations: int = 20, ) -> pv.PolyData: """ Make contours from a labelmap image. + Every label boundary is emitted, including the internal boundaries + between adjacent labels, so the result is a multi-material surface and + is not watertight: an edge where three labels meet is shared by three + faces. Use :meth:`extract_watertight_surface` when a single label's + closed surface is needed. + + Two passes keep the result off the voxel block edges the labelmap is + drawn on. An anisotropic labelmap is first resampled onto an isotropic + grid of its finest pitch, the way :meth:`extract_label_surfaces` does, + so that a boundary between two thick slices lands between them instead + of terracing at one of them. The contour is then Taubin-smoothed, + which the surface net's own constrained smoothing cannot substitute + for: that one may not move a point more than about a voxel, so it + rounds the blocks without removing them. Taubin does not shrink the + surface, and the surface net shares its points between neighboring + labels, so smoothing the mesh as one moves a shared point once and the + labels stay in contact. + Args: labelmap_image (itk.image): The labelmap image to create contours from + smoothing_iterations: Surface-net smoothing iterations. + smoothing_scale: Surface-net smoothing scale. + surface_reduction_rate: Fraction of triangles to remove afterwards + (0.0 disables). + taubin_iterations: Taubin smoothing iterations applied after + reduction, so they act on evenly sized triangles (0 disables). Returns: pv.PolyData: The contours as a PyVista PolyData object """ + spacing = np.asarray(labelmap_image.GetSpacing(), dtype=np.float64) + if float(np.max(spacing) / np.min(spacing)) > _CONTOUR_ANISOTROPY_LIMIT: + self.log_info( + "Contouring on a %.3g mm isotropic grid rather than the " + "labelmap's %s mm one", + float(np.min(spacing)), + " x ".join(f"{value:.3g}" for value in spacing), + ) + labelmap_image = self._resample_labelmap_isotropic( + labelmap_image, float(np.min(spacing)) + ) + labels = pv.wrap(itk.vtk_image_from_image(labelmap_image)) contours = cast( pv.PolyData, @@ -59,42 +169,637 @@ def extract_contours( ), ) - return contours + return self.remesh_and_smooth_surface( + contours, surface_reduction_rate, taubin_iterations + ) + + def extract_watertight_surface( + self, + mask_image: itk.image, + smoothing_iterations: int = 10, + gaussian_sigma_mm: float = 0.5, + surface_reduction_rate: float = 0.0, + anatomy_names: Optional[Sequence[str]] = None, + ) -> pv.PolyData: + """Extract one binary mask's closed, outward-oriented surface. + + :meth:`extract_contours` cannot produce a watertight surface: its + surface nets pinch where a mask self-touches across a voxel diagonal, + leaving edges shared by four faces, and they leave the mask open where + it reaches the image border. Isocontouring a continuous field cannot + pinch, so this pads the mask with one voxel of background, blurs it, and + runs marching cubes at the half-way isovalue instead. + + Reduction goes through :meth:`remesh_and_smooth_surface`, whose ACVD + remeshing keeps a watertight input watertight where the VTK decimators + do not. That is a property of the remesher rather than a guarantee, so + a reduced result is still checked and a warning logged if it degrades. + + Args: + mask_image: Binary mask holding the single structure to extract. + smoothing_iterations: Taubin smoothing iterations (0 disables). + gaussian_sigma_mm: Blur applied before isocontouring, in + millimeters, so it is independent of the voxel pitch. + surface_reduction_rate: Fraction of triangles to remove (0.0 + disables). + anatomy_names: Names passed to :meth:`apply_anatomy_color`, most + specific first. ``None`` leaves the surface uncolored. + + Returns: + The structure's surface, with outward normals. + """ + spacing = np.asarray(mask_image.GetSpacing(), dtype=np.float64) + direction = itk.array_from_matrix(mask_image.GetDirection()) + # One voxel of background all around, so a structure reaching the image + # border still closes; the origin steps back one voxel to match. + padded_arr = np.pad(itk.GetArrayViewFromImage(mask_image).astype(np.float32), 1) + padded = itk.GetImageFromArray(np.ascontiguousarray(padded_arr)) + padded.SetSpacing(mask_image.GetSpacing()) + padded.SetDirection(mask_image.GetDirection()) + padded.SetOrigin( + np.asarray(mask_image.GetOrigin(), dtype=np.float64) + direction @ -spacing + ) + + blurred = itk.smoothing_recursive_gaussian_image_filter( + padded, sigma=gaussian_sigma_mm + ) + surface = cast( + pv.PolyData, + pv.wrap(itk.vtk_image_from_image(blurred)).contour( + [0.5], method="flying_edges" + ), + ) + if smoothing_iterations > 0: + surface = surface.smooth_taubin(n_iter=smoothing_iterations) + # VTK winds faces for a right-handed direction matrix, so an LPS image + # with a negative-determinant direction comes out with inward normals. + surface = surface.compute_normals( + auto_orient_normals=True, consistent_normals=True + ) + + if surface_reduction_rate > 0.0: + surface = self.remesh_and_smooth_surface(surface, surface_reduction_rate, 0) + if not self.is_watertight(surface): + self.log_warning( + "Remeshing by %.2f made the surface non-watertight", + surface_reduction_rate, + ) + + if anatomy_names is not None: + self.apply_anatomy_color(surface, anatomy_names) + return surface + + def extract_label_surfaces( + self, + labelmap_image: itk.image, + isotropic_spacing_mm: Optional[float] = None, + distance_sigma_mm: Optional[float] = None, + smoothing_iterations: int = 30, + ) -> dict[int, pv.PolyData]: + """Extract every label's surface, smooth and conforming with its neighbors. + + :meth:`extract_watertight_surface`, run per label, traces each label's + own voxel block edges: on anisotropic data the result terraces at the + slice pitch, and neighboring labels are contoured independently, so + their shared wall is meshed twice and the two copies do not match. + This extracts all labels together instead: + + 1. The labelmap is resampled onto an isotropic grid with ITK's + label-aware Gaussian interpolator, which votes over the labels in a + physical-space kernel rather than picking a nearest voxel. That is + what removes the terracing: a boundary between two slices lands + between them instead of on one of them. + 2. Every label, plus the background, gets a signed distance map on that + grid, and each voxel is assigned to the label whose map is smallest. + The assignment is a partition, so no gap or overlap can arise. + 3. Label ``L``'s surface is the zero level of ``D_L`` minus the smallest + of the other maps. On a wall between ``L`` and ``M`` that field is + the negation of ``M``'s, so marching cubes puts identical vertices on + both surfaces and the two meet exactly. + 4. The surfaces are merged, which welds those coincident vertices, then + Taubin-smoothed as one mesh. Smoothing therefore moves a shared + vertex once and the surfaces stay in contact. + + Structures thinner than the interpolation kernel lose volume, coronary + arteries most of all; the fraction of each label's voxel volume that + the surface encloses is logged. + + Args: + labelmap_image: Multi-label image; every non-zero label present is + extracted. A binary mask yields a single surface. + isotropic_spacing_mm: Edge length of the isotropic grid the + surfaces are contoured on, which sets both their smoothness and + their triangle count. ``None`` uses the labelmap's finest + spacing. + distance_sigma_mm: Blur applied to the distance maps, which is what + takes the voxel facets out of the contoured surface. ``None`` + uses the isotropic spacing; raising it smooths further and + thins the smallest structures. + smoothing_iterations: Taubin smoothing iterations (0 disables). + + Returns: + Label id → that label's closed, outward-oriented surface. Empty if + the labelmap holds no non-zero label. + """ + labels = itk.GetArrayViewFromImage(labelmap_image) + label_ids = [int(value) for value in np.unique(labels) if value != 0] + if not label_ids: + self.log_warning("Labelmap holds no non-zero label") + return {} + + spacing = np.asarray(labelmap_image.GetSpacing(), dtype=np.float64) + iso = ( + float(np.min(spacing)) + if isotropic_spacing_mm is None + else isotropic_spacing_mm + ) + sigma = iso if distance_sigma_mm is None else distance_sigma_mm + cropped = self._crop_to_labels(labelmap_image, labels) + fine = self._resample_labelmap_isotropic(cropped, iso) + fine_labels = itk.GetArrayViewFromImage(fine) + + # Pass one: the closest and second closest label at every voxel. Two + # values are needed because a label's own map is the closest one inside + # it, and its surface is measured against the next closest. + closest = np.full(fine_labels.shape, np.inf, dtype=np.float32) + runner_up = np.full(fine_labels.shape, np.inf, dtype=np.float32) + closest_index = np.zeros(fine_labels.shape, dtype=np.int16) + for index, label_id in enumerate(label_ids + [0]): + distance = self._signed_distance_mm(fine, fine_labels, label_id, sigma) + is_closer = distance < closest + runner_up = np.where(is_closer, closest, np.minimum(runner_up, distance)) + closest = np.where(is_closer, distance, closest) + closest_index[is_closer] = index + + # Pass two: one surface per label, tagged so the merged mesh can be + # split again after smoothing. The distance maps are recomputed rather + # than kept, which costs one more pass but not one array per label. + parts: list[pv.PolyData] = [] + for index, label_id in enumerate(label_ids): + distance = self._signed_distance_mm(fine, fine_labels, label_id, sigma) + others = np.where(closest_index == index, runner_up, closest) + field = itk.GetImageFromArray(np.ascontiguousarray(distance - others)) + field.CopyInformation(fine) + part = cast( + pv.PolyData, + pv.wrap(itk.vtk_image_from_image(field)).contour( + [0.0], method="flying_edges" + ), + ) + # Voxels where two labels tie exactly contour to zero-area + # triangles. clean turns those into line cells rather than + # dropping them, so the polygons are then taken on their own. + part = part.clean().triangulate() + part = pv.PolyData(part.points, faces=part.faces) + part.cell_data["LabelId"] = np.full(part.n_cells, label_id, dtype=np.int32) + parts.append(part) + + merged = cast(pv.PolyData, pv.merge(parts, merge_points=True)) + if smoothing_iterations > 0: + # Every wall between two labels is meshed twice, so its edges are + # non-manifold; without non_manifold_smoothing VTK pins them and + # nothing moves. + merged = merged.smooth_taubin( + n_iter=smoothing_iterations, non_manifold_smoothing=True + ) + + voxel_volume = float(np.prod(spacing)) + merged_ids = np.asarray(merged.cell_data["LabelId"]) + surfaces: dict[int, pv.PolyData] = {} + for label_id in label_ids: + cell_ids: list[int] = np.flatnonzero(merged_ids == label_id).tolist() + surface = self.extract_surface(merged.extract_cells(cell_ids)).triangulate() + # The bookkeeping arrays of the merge and the split; the label is + # the key of the returned mapping, so it is not data on the mesh. + for array_name in ("LabelId", "vtkOriginalCellIds", "vtkOriginalPointIds"): + surface.cell_data.pop(array_name, None) + surface.point_data.pop(array_name, None) + # VTK winds faces for a right-handed direction matrix, so an LPS + # image with a negative-determinant direction comes out inward. + surfaces[label_id] = surface.compute_normals( + auto_orient_normals=True, consistent_normals=True + ) + self.log_debug( + "Label %d: %d triangles, %.3f of its voxel volume", + label_id, + surfaces[label_id].n_cells, + float(surfaces[label_id].volume) + / (int(np.count_nonzero(labels == label_id)) * voxel_volume), + ) + return surfaces + + @staticmethod + def _crop_to_labels(labelmap_image: itk.image, labels: np.ndarray) -> itk.image: + """Return *labelmap_image* cropped to its labels and padded with background. + + The pad closes structures that reach the image border, which would + otherwise contour to an open surface, and gives the distance maps room + to fall away from the outermost label. + """ + spacing = np.asarray(labelmap_image.GetSpacing(), dtype=np.float64) + direction = itk.array_from_matrix(labelmap_image.GetDirection()) + # labels' axes are reversed relative to the ITK image, so the extents + # come back as (z, y, x) and are flipped to (x, y, z) for the origin. + extents = np.nonzero(labels) + start = np.array([int(axis.min()) for axis in extents]) + stop = np.array([int(axis.max()) + 1 for axis in extents]) + cropped_arr = np.pad( + labels[start[0] : stop[0], start[1] : stop[1], start[2] : stop[2]], + _LABEL_SURFACE_PAD, + ) + cropped = itk.GetImageFromArray( + np.ascontiguousarray(cropped_arr.astype(np.uint16)) + ) + cropped.SetSpacing(labelmap_image.GetSpacing()) + cropped.SetDirection(labelmap_image.GetDirection()) + cropped.SetOrigin( + np.asarray(labelmap_image.GetOrigin(), dtype=np.float64) + + direction @ (spacing * (start[::-1] - _LABEL_SURFACE_PAD)) + ) + return cropped + + @staticmethod + def _resample_labelmap_isotropic( + labelmap_image: itk.image, + isotropic_spacing_mm: float, + sigma_mm: Optional[np.ndarray] = None, + ) -> itk.image: + """Resample a labelmap onto an isotropic grid, label boundaries intact. + + ``LabelImageGaussianInterpolateImageFunction`` gives each output voxel + the label with the largest Gaussian-weighted vote among its neighbors, + so labels stay whole numbers, keep sharing their walls, and their + boundaries move to where the vote turns over rather than snapping to an + input voxel. + + Args: + labelmap_image: Labelmap to resample. + isotropic_spacing_mm: Edge length of the output voxels. + sigma_mm: Per-axis width of the voting kernel. ``None`` uses one + input voxel along each axis, which is what lets a coarse axis + interpolate between its slices. + """ + spacing = np.asarray(labelmap_image.GetSpacing(), dtype=np.float64) + size = itk.size(labelmap_image) + interpolator = itk.LabelImageGaussianInterpolateImageFunction.New( + labelmap_image + ) + sigma = spacing if sigma_mm is None else sigma_mm + interpolator.SetSigma([float(value) for value in sigma]) + interpolator.SetAlpha(3.0) + return itk.resample_image_filter( + labelmap_image, + size=[ + int((size[axis] - 1) * spacing[axis] / isotropic_spacing_mm) + 1 + for axis in range(3) + ], + output_spacing=[isotropic_spacing_mm] * 3, + output_origin=list(labelmap_image.GetOrigin()), + output_direction=labelmap_image.GetDirection(), + interpolator=interpolator, + default_pixel_value=0, + ) + + @staticmethod + def _signed_distance_mm( + reference: itk.image, labels: np.ndarray, label_id: int, sigma_mm: float + ) -> np.ndarray: + """Return the signed distance, in mm, to *label_id*, negative inside. + + The map measures to voxel centers, so its zero level is faceted at the + voxel pitch; *sigma_mm* of blur takes those facets out, at the cost of + pulling the level set in by roughly ``sigma_mm ** 2`` times the surface + curvature. ``label_id`` of ``0`` measures to the background, whose + distance is the negated distance to the union of every label. + """ + inside = labels != 0 if label_id == 0 else labels == label_id + mask = itk.GetImageFromArray(np.ascontiguousarray(inside.astype(np.uint8))) + mask.CopyInformation(reference) + distance = itk.signed_maurer_distance_map_image_filter( + mask, + use_image_spacing=True, + squared_distance=False, + inside_is_positive=False, + ) + if sigma_mm > 0.0: + distance = itk.smoothing_recursive_gaussian_image_filter( + distance, sigma=sigma_mm + ) + distance_arr = np.asarray(itk.GetArrayFromImage(distance), dtype=np.float32) + return -distance_arr if label_id == 0 else distance_arr @staticmethod - def smooth_and_decimate_surface( + def is_watertight(surface: pv.PolyData) -> bool: + """Report whether every edge of *surface* is shared by exactly two faces.""" + faces = surface.triangulate().faces.reshape(-1, 4)[:, 1:] + edges = np.sort( + np.vstack([faces[:, [0, 1]], faces[:, [1, 2]], faces[:, [2, 0]]]), axis=1 + ) + _, counts = np.unique(edges, axis=0, return_counts=True) + return bool(np.all(counts == 2)) + + def extract_tetrahedra( + self, + mask_image: itk.image, + element_size_mm: Optional[float] = None, + anatomy_names: Optional[Sequence[str]] = None, + ) -> pv.UnstructuredGrid: + """Build a tetrahedral mesh filling one binary mask. + + Every retained voxel becomes a hexahedron, which VTK then splits into + six tetrahedra sharing the hexahedra's points, so the result is a + conforming mesh whose boundary is the voxel staircase rather than the + smooth surface :meth:`extract_label_surfaces` returns. Pass the result + through :meth:`trim_tetrahedra_to_surface` to relax that staircase onto + the surface. + + Args: + mask_image: Binary mask holding the single structure to fill. + element_size_mm: Edge length of the isotropic voxels the mask is + resampled to before meshing, which is the resulting element + size. ``None`` meshes the mask's own voxels, so on anisotropic + data the elements inherit that anisotropy. A size above the + thinnest part of the structure drops that part. + anatomy_names: Names passed to :meth:`apply_anatomy_color`, most + specific first. ``None`` leaves the mesh uncolored. + + Returns: + The structure's tetrahedral mesh, empty if *element_size_mm* is too + coarse to keep any of the mask. + """ + mask_arr = itk.GetArrayViewFromImage(mask_image) != 0 + # mask_arr axes are reversed relative to the ITK image, so the per-axis + # extents come back as (z, y, x) and are flipped to (x, y, z). + starts, stops = [], [] + for axis_extent in np.nonzero(mask_arr): + starts.append(int(axis_extent.min())) + stops.append(int(axis_extent.max()) + 1) + start_zyx, stop_zyx = np.array(starts), np.array(stops) + cropped_arr = mask_arr[ + start_zyx[0] : stop_zyx[0], + start_zyx[1] : stop_zyx[1], + start_zyx[2] : stop_zyx[2], + ] + + spacing = np.asarray(mask_image.GetSpacing(), dtype=np.float64) + direction = itk.array_from_matrix(mask_image.GetDirection()) + # Cropping only translates the image, so the direction is unchanged and + # the winding correction below still applies after any resampling. + cropped = itk.GetImageFromArray( + np.ascontiguousarray(cropped_arr.astype(np.uint16)) + ) + cropped.SetSpacing(mask_image.GetSpacing()) + cropped.SetDirection(mask_image.GetDirection()) + cropped.SetOrigin( + np.asarray(mask_image.GetOrigin(), dtype=np.float64) + + direction @ (spacing * start_zyx[::-1]) + ) + + if element_size_mm is not None: + # A vote over the mask rather than a nearest neighbor, and one + # taken over at least an output voxel, so that coarsening keeps + # thin walls instead of sampling through them. + cropped = self._resample_labelmap_isotropic( + cropped, + element_size_mm, + np.maximum(spacing, element_size_mm), + ) + cropped_arr = itk.GetArrayViewFromImage(cropped) != 0 + spacing = np.full(3, element_size_mm, dtype=np.float64) + if not cropped_arr.any(): + self.log_warning( + "Elements of %.3g mm are too coarse for this structure; " + "its mesh is empty", + element_size_mm, + ) + return pv.UnstructuredGrid() + + # Corner-point grid: one more point than voxels along each axis, with + # the origin backed off half a voxel to reach the first voxel's corner. + grid = pv.ImageData( + dimensions=tuple(int(n) + 1 for n in cropped_arr.shape[::-1]), + spacing=tuple(spacing), + origin=tuple( + np.asarray(cropped.GetOrigin(), dtype=np.float64) + + direction @ (spacing * -0.5) + ), + ) + grid.direction_matrix = direction + grid.cell_data["label"] = cropped_arr.ravel().astype(np.uint8) + voxels = grid.threshold(0.5, scalars="label") + + order = ( + _VOXEL_TO_HEX if np.linalg.det(direction) > 0.0 else _VOXEL_TO_HEX_FLIPPED + ) + connectivity = voxels.cells.reshape(-1, 9)[:, 1:][:, order] + cells = np.hstack( + [np.full((len(connectivity), 1), 8, dtype=connectivity.dtype), connectivity] + ) + hexahedra = pv.UnstructuredGrid( + cells.ravel(), + np.full(len(connectivity), pv.CellType.HEXAHEDRON, dtype=np.uint8), + voxels.points, + ) + tetrahedra = cast(pv.UnstructuredGrid, hexahedra.triangulate()) + + if anatomy_names is not None: + self.apply_anatomy_color(tetrahedra, anatomy_names) + return tetrahedra + + def trim_tetrahedra_to_surface( + self, + tetrahedra: pv.UnstructuredGrid, surface: pv.PolyData, - decimation_reduction: float, - smoothing_iterations: int, - ) -> pv.PolyData: - """Optionally decimate then smooth a surface (no-op when disabled). + iterations: int = 5, + relaxation: float = 0.6, + min_scaled_jacobian: float = 0.1, + ) -> pv.UnstructuredGrid: + """Relax a tetrahedral mesh onto *surface*, keeping every cell whole. + + :meth:`extract_tetrahedra` meshes voxels, so its boundary is a + staircase that both protrudes through the smooth surface + :meth:`extract_label_surfaces` builds from the same mask and falls short + of it elsewhere. Cutting the mesh at the surface with ``clip_surface`` + would follow it exactly but shatters the boundary tetrahedra into + slivers (about a tenth of the cells drop below a scaled Jacobian of + 0.1), and neither VTK nor any current dependency can repair those. + + So nothing is cut. A crinkle clip drops the cells that lie entirely + outside while leaving every surviving cell intact, then the mesh is + relaxed: each sweep moves every point part of the way toward the + average of its neighbors, and every boundary point instead toward its + closest point on *surface*. The interior smoothing is what makes room + for the boundary to reach the surface -- projecting the boundary alone + flattens the cells behind it, and the quality bound below then undoes + the move, which is why one projection pass leaves the staircase in + place. + + A move that would still wreck a cell is backed off: every point of a + cell below *min_scaled_jacobian* has that sweep's step halved, + repeatedly, until the whole mesh clears the bound. + + Args: + tetrahedra: Volume mesh to relax; its cell and field data survive. + surface: Closed surface to relax onto, in the same frame. + iterations: Relaxation sweeps. The boundary reaches the surface in + the first few and then stops moving: on the Duke heart labels, + sweeps beyond the default leave the mean boundary-to-surface + distance and the worst cell quality where they already were, + and only cost time. + relaxation: Fraction of the way to its target a point moves per + sweep. ``1.0`` moves the whole way and oscillates. + min_scaled_jacobian: Cell-quality bound every tetrahedron must meet + after each sweep. ``0.0`` only rules out flattened and + inverted cells; the default also rules out slivers. + + Returns: + The relaxed mesh. + """ + # crinkle keeps whole cells, so this only discards, never subdivides. + relaxed = cast( + pv.UnstructuredGrid, + tetrahedra.clip_surface(surface, invert=True, crinkle=True), + ) + if relaxed.n_cells == 0: + self.log_warning("Trimming against the surface removed every cell") + return relaxed + + connectivity = relaxed.cells_dict[np.uint8(pv.CellType.TETRA)] + # Both directions of every tetrahedron edge, so a point's neighbors are + # the second column of the rows its id occupies in the first. + edges = np.vstack( + [ + connectivity[:, pair] + for pair in ((0, 1), (0, 2), (0, 3), (1, 2), (1, 3), (2, 3)) + ] + ) + starts = np.concatenate([edges[:, 0], edges[:, 1]]) + ends = np.concatenate([edges[:, 1], edges[:, 0]]) + point_count = relaxed.n_points + neighbor_counts = np.bincount(starts, minlength=point_count).clip(1) + boundary_ids = np.asarray( + relaxed.extract_surface(algorithm="dataset_surface").point_data[ + "vtkOriginalPointIds" + ] + ) - Decimation uses ``decimate_pro`` on a triangulated copy; because - ``decimate_pro`` discards cell data, per-cell ``boundary_labels`` (needed - for anatomy splitting downstream) are transferred back onto the decimated - cells from their nearest original cell so anatomy materials still apply. - Smoothing uses non-shrinking Taubin smoothing, which only moves points and - therefore preserves cells and their labels. + points = np.array(relaxed.points) + for _ in range(iterations): + target = ( + np.column_stack( + [ + np.bincount( + starts, weights=points[ends, axis], minlength=point_count + ) + for axis in range(3) + ] + ) + / neighbor_counts[:, np.newaxis] + ) + _, closest = cast( + "tuple[np.ndarray, np.ndarray]", + surface.find_closest_cell( + target[boundary_ids], return_closest_point=True + ), + ) + target[boundary_ids] = closest + step = relaxation * (target - points) + + # Backing one point off can push a neighboring cell below the bound + # that the full move had cleared, so this repeats until the whole + # mesh passes. The step bottoms out at zero, so the loop + # terminates even when the mesh holds cells below the bound that no + # amount of backing off can rescue. + damping = np.ones(point_count) + while True: + relaxed.points = points + damping[:, np.newaxis] * step + quality = np.asarray( + relaxed.cell_quality(["scaled_jacobian"]).cell_data[ + "scaled_jacobian" + ] + ) + below_bound = quality < min_scaled_jacobian + if not np.any(below_bound): + break + damped_ids = np.unique(connectivity[below_bound]) + damped_ids = damped_ids[damping[damped_ids] > 0.0] + if damped_ids.size == 0: + break + halved = damping[damped_ids] * 0.5 + damping[damped_ids] = np.where(halved < _MIN_TRIM_DAMPING, 0.0, halved) + points = np.array(relaxed.points) + return relaxed + + def remesh_and_smooth_surface( + self, + surface: pv.PolyData, + surface_reduction_rate: float = 0.0, + smoothing_iterations: int = 0, + ) -> pv.PolyData: + """Optionally remesh then smooth a surface (no-op when disabled). + + Reduction is isotropic remeshing (ACVD, via ``pyacvd``) rather than + decimation: the surface is re-tiled with uniform, well-shaped triangles + at the requested resolution. ``decimate_pro`` reaches the same triangle + count but leaves a watertight input non-watertight; ACVD does not. + + Remeshing rebuilds the topology and so discards cell data, exactly as + ``decimate_pro`` did: per-cell ``boundary_labels`` (needed for anatomy + splitting downstream) are transferred back onto the new cells from their + nearest original cell so anatomy materials still apply. Uniform + triangles cannot represent a label patch smaller than one of them, + though, so such a patch is absorbed by its neighbours and its label pair + disappears -- a warning names the pairs lost. ``decimate_pro`` kept + those patches by being non-uniform, which is the trade being made here. + Smoothing uses non-shrinking Taubin smoothing, which only moves points + and therefore preserves cells and their labels. It is told to move + non-manifold points too, since on a multi-material surface every edge + where three labels meet is non-manifold and VTK pins those points + otherwise; on a manifold surface the setting has nothing to act on. Args: surface: Input surface. - decimation_reduction: Fraction of triangles to remove (0.0 disables). + surface_reduction_rate: Fraction of triangles to remove (0.0 disables). smoothing_iterations: Taubin smoothing iterations (0 disables). Returns: - The decimated and/or smoothed surface. + The remeshed and/or smoothed surface. """ conditioned = surface - if decimation_reduction > 0.0: + if surface_reduction_rate > 0.0: original = conditioned - conditioned = conditioned.triangulate().decimate_pro(decimation_reduction) + clustering = pyacvd.Clustering(conditioned.triangulate()) + # One cluster per retained point. A closed surface carries about + # twice as many triangles as points, so scaling the point count by + # (1 - rate) scales the triangle count by the same fraction; four + # is the fewest clusters that can still close a surface. + clustering.cluster( + max(4, round(original.n_points * (1.0 - surface_reduction_rate))) + ) + conditioned = clustering.create_mesh() if "boundary_labels" in original.cell_data: + labels = np.asarray(original.cell_data["boundary_labels"]) nearest = original.find_closest_cell(conditioned.cell_centers().points) - conditioned.cell_data["boundary_labels"] = np.asarray( - original.cell_data["boundary_labels"] - )[nearest] + conditioned.cell_data["boundary_labels"] = labels[nearest] + + pairs = labels.reshape(len(labels), -1) + before = {tuple(row) for row in np.unique(pairs, axis=0).tolist()} + after = { + tuple(row) for row in np.unique(pairs[nearest], axis=0).tolist() + } + if before - after: + self.log_warning( + "Remeshing by %.2f dropped %d of %d boundary label pairs, " + "each covering less than one output triangle: %s", + surface_reduction_rate, + len(before - after), + len(before), + sorted(before - after), + ) if smoothing_iterations > 0: - conditioned = conditioned.smooth_taubin(n_iter=smoothing_iterations) + conditioned = conditioned.smooth_taubin( + n_iter=smoothing_iterations, non_manifold_smoothing=True + ) return conditioned @staticmethod diff --git a/src/physiotwin4d/convert_vtk_to_usd.py b/src/physiotwin4d/convert_vtk_to_usd.py index 6ccc1831..3464ac0b 100644 --- a/src/physiotwin4d/convert_vtk_to_usd.py +++ b/src/physiotwin4d/convert_vtk_to_usd.py @@ -98,7 +98,10 @@ def __init__( input_polydata: Sequence of PyVista/VTK meshes (one per time step, or one per static object when static_merge is True) mask_ids: Optional mapping of label IDs to anatomical region names. - If provided, meshes will be split by labeled regions. + If provided, meshes are split by labeled region, one prim + per region, and the static_merge layout is not used: with + a single mesh that yields one static prim per structure, + with several it yields one time-varying prim per structure. compute_normals: Whether to compute vertex normals convert_to_surface: If True, extract surface from volumetric meshes frames_per_second: Time codes per second (default 24.0). @@ -153,6 +156,13 @@ def __init__( self.colormap: str = "plasma" self.intensity_range: Optional[tuple[float, float]] = None + if static_merge and mask_ids and len(self.input_polydata) > 1: + raise ValueError( + "static_merge with mask_ids is only defined for a single mesh: " + "several static objects each holding the same labels would " + "collide on one prim path per label. Merge them first, or drop " + "static_merge to treat them as frames." + ) if not static_merge and time_codes is not None: if len(time_codes) != len(self.input_polydata): raise ValueError( @@ -649,12 +659,14 @@ def convert( material_mgr = MaterialManager(stage) mesh_converter = UsdMeshConverter(stage, self.settings, material_mgr) - # Process meshes - if self._is_static_merge: - self._convert_static_merge(stage, root_path, material_mgr, mesh_converter) - elif self.mask_ids: + # Process meshes. Labels win over the static layout: a per-cell label + # array names the structures outright, which the static layout's one + # prim per input mesh cannot. + if self.mask_ids: # Split by anatomical regions self._convert_with_labels(stage, root_path, material_mgr, mesh_converter) + elif self._is_static_merge: + self._convert_static_merge(stage, root_path, material_mgr, mesh_converter) else: # Single mesh (or time series) conversion self._convert_unified(stage, root_path, material_mgr, mesh_converter) diff --git a/src/physiotwin4d/image_tools.py b/src/physiotwin4d/image_tools.py index 284c71fb..c0b4b7c8 100644 --- a/src/physiotwin4d/image_tools.py +++ b/src/physiotwin4d/image_tools.py @@ -270,6 +270,65 @@ def make_isotropic_image(self, image: itk.Image) -> itk.Image: result.DisconnectPipeline() return result + def resample_image_by_scale( + self, image: itk.Image, scale: float, interpolate: bool = True + ) -> itk.Image: + """Resample a 3-D *image* to *scale* times its voxel count per axis. + + The physical extent is preserved: spacing is rescaled to compensate for + the new voxel count, and the origin shifts by half the spacing change so + the resampled voxel centers stay inside the original extent. + + Args: + image: 3-D ITK image to resample. + scale: Per-axis voxel-count multiplier. Values below ``1.0`` + coarsen, values above ``1.0`` upsample. + interpolate: Use linear interpolation. ``False`` selects nearest + neighbor, which is what labelmaps need. + + Returns: + Resampled image, covering the same physical extent as *image*. + + Raises: + ValueError: If *image* is not 3-D, or *scale* is not positive. + """ + if image.GetImageDimension() != 3: + raise ValueError( + f"resample_image_by_scale requires a 3-D image; " + f"got {image.GetImageDimension()}-D" + ) + if scale <= 0.0: + raise ValueError(f"scale must be positive; got {scale}") + + spacing = np.asarray(image.GetSpacing(), dtype=np.float64) + size = np.asarray(image.GetLargestPossibleRegion().GetSize(), dtype=np.int64) + new_size = np.maximum(1, np.ceil(size * scale)).astype(np.int64) + new_spacing = spacing * size / new_size + + ImageType = type(image) + if interpolate: + interpolator = itk.LinearInterpolateImageFunction[ImageType, itk.D].New() + else: + interpolator = itk.NearestNeighborInterpolateImageFunction[ + ImageType, itk.D + ].New() + + direction = itk.array_from_matrix(image.GetDirection()) + resampler = itk.ResampleImageFilter[ImageType, ImageType].New() + resampler.SetInput(image) + resampler.SetInterpolator(interpolator) + resampler.SetOutputSpacing([float(v) for v in new_spacing]) + resampler.SetSize([int(n) for n in new_size]) + resampler.SetOutputOrigin( + np.asarray(image.GetOrigin(), dtype=np.float64) + + direction @ ((new_spacing - spacing) / 2.0) + ) + resampler.SetOutputDirection(image.GetDirection()) + resampler.Update() + result = resampler.GetOutput() + result.DisconnectPipeline() + return result + @staticmethod def _per_axis_values( value: Union[float, int, list, tuple, NDArray[Any]], diff --git a/src/physiotwin4d/register_images_base.py b/src/physiotwin4d/register_images_base.py index c148899c..92f82693 100644 --- a/src/physiotwin4d/register_images_base.py +++ b/src/physiotwin4d/register_images_base.py @@ -667,12 +667,19 @@ def _capture_delegate_result( def get_registered_image(self) -> itk.Image: """Get the registered image. + The moving image is an intensity image, so voxels of the fixed grid that + fall outside it are filled with :meth:`_prewarp_background_value` rather + than 0, which is a tissue intensity rather than an absence of tissue. + Returns: itk.Image: The registered image """ if self.moving_image_registered is None: TfmTools = TransformTools() self.moving_image_registered = TfmTools.transform_image( - self.moving_image, self.forward_transform, self.fixed_image + self.moving_image, + self.forward_transform, + self.fixed_image, + background_value=self._prewarp_background_value(self.moving_image), ) return self.moving_image_registered diff --git a/src/physiotwin4d/register_time_series_images.py b/src/physiotwin4d/register_time_series_images.py index 16a481ac..c698d7fc 100644 --- a/src/physiotwin4d/register_time_series_images.py +++ b/src/physiotwin4d/register_time_series_images.py @@ -385,9 +385,14 @@ def reconstruct_time_series( # Use fixed image as reference reference_image = moving_image - # Transform the moving image to the reference space + # Transform the moving image to the reference space. The fixed + # image is an intensity image, so voxels sampled outside it take the + # modality's "no tissue" value, not 0. reconstructed = self.transform_tools.transform_image( - self.fixed_image, inverse_transform, reference_image + self.fixed_image, + inverse_transform, + reference_image, + background_value=self._prewarp_background_value(self.fixed_image), ) reconstructed_images.append(reconstructed) diff --git a/src/physiotwin4d/transform_tools.py b/src/physiotwin4d/transform_tools.py index 5f12bed2..49e3aeb4 100644 --- a/src/physiotwin4d/transform_tools.py +++ b/src/physiotwin4d/transform_tools.py @@ -518,13 +518,23 @@ def transform_image( ) # ITK's wrapping types DefaultPixelValue to the image's pixel type, and - # rejects a Python float for an integer image. + # rejects a Python float for a discrete image. dtype = itk.GetArrayViewFromImage(img).dtype - default_pixel_value: Union[int, float] = ( - int(round(background_value)) - if np.issubdtype(dtype, np.integer) - else float(background_value) - ) + default_pixel_value: Union[int, float] + if np.issubdtype(dtype, np.integer) or np.issubdtype(dtype, np.bool_): + default_pixel_value = int(round(background_value)) + low, high = ( + (0, 1) + if np.issubdtype(dtype, np.bool_) + else (int(np.iinfo(dtype).min), int(np.iinfo(dtype).max)) + ) + if not low <= default_pixel_value <= high: + raise ValueError( + f"background_value {background_value} is outside the range " + f"[{low}, {high}] of the image's {dtype} pixel type" + ) + else: + default_pixel_value = float(background_value) img_reg = itk.resample_image_filter( Input=img, diff --git a/src/physiotwin4d/usd_tools.py b/src/physiotwin4d/usd_tools.py index 51617014..a5b22988 100644 --- a/src/physiotwin4d/usd_tools.py +++ b/src/physiotwin4d/usd_tools.py @@ -1276,7 +1276,11 @@ def list_mesh_paths_under( self, stage_or_path: Usd.Stage | str, parent_path: str = "/World/Meshes" ) -> list[str]: """ - List paths of all mesh prims under a parent path. + List paths of all mesh prims at any depth under a parent path. + + Descends the whole subtree, so meshes written under an intermediate + Xform - as :class:`ConvertVTKToUSD` writes labeled structures, at + ``/World/{basename}/{anatomy_group}/{structure}`` - are found too. Args: stage_or_path: USD Stage or path to USD file @@ -1293,11 +1297,11 @@ def list_mesh_paths_under( parent = stage.GetPrimAtPath(parent_path) if not parent.IsValid(): return [] - result = [] - for prim in parent.GetAllChildren(): - if prim.IsA(UsdGeom.Mesh): - result.append(str(prim.GetPath())) - return result + return [ + str(prim.GetPath()) + for prim in Usd.PrimRange(parent) + if prim.IsA(UsdGeom.Mesh) + ] def repair_mesh_primvar_element_sizes( self, diff --git a/src/physiotwin4d/vtk_to_usd/usd_mesh_converter.py b/src/physiotwin4d/vtk_to_usd/usd_mesh_converter.py index f0c5d135..0274c2db 100644 --- a/src/physiotwin4d/vtk_to_usd/usd_mesh_converter.py +++ b/src/physiotwin4d/vtk_to_usd/usd_mesh_converter.py @@ -49,10 +49,31 @@ def __init__( self.stage = stage self.settings = settings self.material_mgr = material_mgr - # Most recent triangulation face-map from create_mesh(). Reused by - # create_time_varying_mesh() when writing per-time-step cell primvars, - # since topology (and thus the map) is invariant across time samples. - self._last_triangulation_face_map: Optional[np.ndarray] = None + + def _resolve_topology( + self, mesh_data: MeshData + ) -> tuple[np.ndarray, np.ndarray, Optional[np.ndarray]]: + """Return the face counts, indices and triangulation map to author. + + Triangulation only happens when it was asked for and the mesh holds + faces that are not triangles; the map is ``None`` otherwise. + + Args: + mesh_data: Mesh whose topology is being written. + + Returns: + ``(face_vertex_counts, face_vertex_indices, + triangulation_face_map)``, the last mapping each triangulated face + back to its source face. + """ + face_counts = mesh_data.face_vertex_counts + face_indices = mesh_data.face_vertex_indices + if self.settings.triangulate_meshes and not all( + count == 3 for count in face_counts + ): + logger.debug("Triangulating mesh faces") + return triangulate_face(face_counts, face_indices) + return face_counts, face_indices, None def create_mesh( self, @@ -81,26 +102,16 @@ def create_mesh( usd_points = lps_points_to_usd(mesh_data.points) # Handle triangulation if requested - face_counts = mesh_data.face_vertex_counts - face_indices = mesh_data.face_vertex_indices - - triangulation_face_map: Optional[np.ndarray] = None - if self.settings.triangulate_meshes: - # Check if any faces are not triangles - if not all(count == 3 for count in face_counts): - logger.debug("Triangulating mesh faces") - ( - face_counts, - face_indices, - triangulation_face_map, - ) = triangulate_face(face_counts, face_indices) - self._last_triangulation_face_map = triangulation_face_map + face_counts, face_indices, triangulation_face_map = self._resolve_topology( + mesh_data + ) # Convert to Vt arrays face_counts_vt = Vt.IntArray(face_counts.tolist()) face_indices_vt = Vt.IntArray(face_indices.tolist()) - # Set topology (static - doesn't change with time) + # Set topology as the default value. create_time_varying_mesh() adds + # time samples on top of this when a series changes topology. mesh.CreateFaceVertexCountsAttr(face_counts_vt) mesh.CreateFaceVertexIndicesAttr(face_indices_vt) @@ -335,7 +346,14 @@ def create_time_varying_mesh( ) -> UsdGeom.Mesh: """Create a mesh with time-varying attributes. - Assumes constant topology (same number of points/faces). + A series whose frames share one topology, as a surface propagated + through a deformation does, authors that topology once and time-samples + only the point positions, so viewers interpolate between samples. A + series whose frames were built independently, and so agree on neither + point count nor triangulation, additionally time-samples + ``faceVertexCounts`` and ``faceVertexIndices``; USD holds those samples + rather than interpolating them, so such a mesh snaps from frame to + frame. Args: mesh_data_sequence: List of MeshData for each time step @@ -360,15 +378,40 @@ def create_time_varying_mesh( f"with {len(time_codes)} time steps" ) + topologies = [self._resolve_topology(md) for md in mesh_data_sequence] + first_counts, first_indices, _ = topologies[0] + topology_varies = any( + not np.array_equal(counts, first_counts) + or not np.array_equal(indices, first_indices) + for counts, indices, _ in topologies[1:] + ) + # Create mesh with first time step first_mesh_data = mesh_data_sequence[0] mesh = self.create_mesh( first_mesh_data, mesh_path, time_codes[0], bind_material=bind_material ) + if topology_varies: + logger.warning( + "Topology changes across the %d frames of %s; authoring it per " + "time sample, which viewers hold rather than interpolate", + len(time_codes), + mesh_path, + ) + # A time sample wins over the default at every time, so the first + # frame has to be sampled too or it would resolve to the last one. + counts_attr = mesh.GetFaceVertexCountsAttr() + indices_attr = mesh.GetFaceVertexIndicesAttr() + for (counts, indices, _), time_code in zip( + topologies, time_codes, strict=False + ): + counts_attr.Set(Vt.IntArray(counts.tolist()), time_code) + indices_attr.Set(Vt.IntArray(indices.tolist()), time_code) + # Add time samples for subsequent steps - for mesh_data, time_code in zip( - mesh_data_sequence[1:], time_codes[1:], strict=False + for frame_index, (mesh_data, time_code) in enumerate( + zip(mesh_data_sequence[1:], time_codes[1:], strict=False), start=1 ): # Update points usd_points = lps_points_to_usd(mesh_data.points) @@ -387,8 +430,8 @@ def create_time_varying_mesh( if mesh_data.colors is not None: self._add_vertex_colors(mesh, mesh_data.colors, time_code) - # Update generic arrays (reuse the triangulation map computed - # for the first time sample; topology is invariant across time). + # Update generic arrays, expanding uniform ones with this frame's + # own triangulation map. if ( self.settings.preserve_point_arrays or self.settings.preserve_cell_arrays @@ -397,7 +440,7 @@ def create_time_varying_mesh( mesh, mesh_data, time_code, - self._last_triangulation_face_map, + topologies[frame_index][2], ) logger.info(f"Created time-varying mesh with {len(time_codes)} time samples") diff --git a/src/physiotwin4d/workflow_convert_image_to_usd.py b/src/physiotwin4d/workflow_convert_image_to_usd.py index ca639889..a0c32a21 100644 --- a/src/physiotwin4d/workflow_convert_image_to_usd.py +++ b/src/physiotwin4d/workflow_convert_image_to_usd.py @@ -55,6 +55,7 @@ def __init__( dynamic_labelmap_ids: Optional[list[int]] = None, mask_dilation_radius: int = 10, frames_per_second: float = 1.0, + surface_reduction_rate: float = 0.0, log_level: int | str = logging.INFO, save_assets: bool = True, ) -> None: @@ -82,6 +83,9 @@ def __init__( the dynamic/static registration masks. Defaults to 10. frames_per_second: Frames per second for animated USD time series. Defaults to 24.0, matching the underlying VTK-to-USD converter. + surface_reduction_rate: Fraction in ``[0, 1)`` of triangles to + remove from every extracted surface. ``0.0`` (default) skips + decimation. log_level: Logging level (default: logging.INFO) save_assets: Write registered images, transforms, and labelmaps output_directory when True @@ -99,6 +103,7 @@ def __init__( self.dynamic_labelmap_ids = dynamic_labelmap_ids if dynamic_labelmap_ids else [] self.output_directory = output_directory self.frames_per_second = frames_per_second + self.surface_reduction_rate = surface_reduction_rate self.save_assets = save_assets self.registration_results: list[ @@ -355,7 +360,9 @@ def _generate_reference_contours(self) -> None: labelmap = self.reference_segmentation_results["labelmap"] # Generate all anatomy contours - all_contours = self.contour_tools.extract_contours(labelmap) + all_contours = self.contour_tools.extract_contours( + labelmap, surface_reduction_rate=self.surface_reduction_rate + ) self.reference_contours = { "all": all_contours, } @@ -369,7 +376,9 @@ def _generate_reference_contours(self) -> None: ) dynamic_labelmap = itk.GetImageFromArray(dynamic_labelmap_arr) dynamic_labelmap.CopyInformation(labelmap) - dynamic_contours = self.contour_tools.extract_contours(dynamic_labelmap) + dynamic_contours = self.contour_tools.extract_contours( + dynamic_labelmap, surface_reduction_rate=self.surface_reduction_rate + ) static_labelmap_arr = itk.GetArrayFromImage(labelmap) static_labelmap_arr = np.where( @@ -379,7 +388,9 @@ def _generate_reference_contours(self) -> None: ) static_labelmap = itk.GetImageFromArray(static_labelmap_arr) static_labelmap.CopyInformation(labelmap) - static_contours = self.contour_tools.extract_contours(static_labelmap) + static_contours = self.contour_tools.extract_contours( + static_labelmap, surface_reduction_rate=self.surface_reduction_rate + ) # Store reference contours self.reference_contours["dynamic"] = dynamic_contours diff --git a/src/physiotwin4d/workflow_convert_image_to_vtk.py b/src/physiotwin4d/workflow_convert_image_to_vtk.py index 2e06526d..30fb392e 100644 --- a/src/physiotwin4d/workflow_convert_image_to_vtk.py +++ b/src/physiotwin4d/workflow_convert_image_to_vtk.py @@ -17,7 +17,7 @@ ct = itk.imread("chest_ct.nii.gz") segmenter = SegmentChestTotalSegmentatorWithContrast() workflow = WorkflowConvertImageToVTK(segmentation_method=segmenter) - result = workflow.process(ct, surface_target_reduction=0.5) + result = workflow.process(ct, surface_reduction_rate=0.5) # Combined single-file output (default) ContourTools.save_combined_surfaces(result["surfaces"], "./out/patient.vtp") @@ -222,7 +222,7 @@ def process( self, input_image: Any, anatomy_groups: Optional[list[str]] = None, - surface_target_reduction: float = 0.0, + surface_reduction_rate: float = 0.0, extract_label_surfaces: bool = False, ) -> dict[str, Any]: """Segment the CT image and extract per-anatomy-group VTK surfaces. @@ -233,8 +233,8 @@ def process( processes all non-empty groups. Valid names are given by :attr:`anatomy_groups`, derived from the active segmenter's taxonomy. - surface_target_reduction: Fraction in ``[0, 1)`` of surface - triangles to remove via ``decimate_pro(surface_target_reduction, + surface_reduction_rate: Fraction in ``[0, 1)`` of surface + triangles to remove via ``decimate_pro(surface_reduction_rate, preserve_topology=True)``. ``0.0`` (default) skips decimation. Applied to both group and (when requested) label surfaces. extract_label_surfaces: When ``True``, also extract one surface per @@ -313,9 +313,9 @@ def process( continue export_surface = base_surface - if surface_target_reduction > 0.0: + if surface_reduction_rate > 0.0: export_surface = export_surface.decimate_pro( - surface_target_reduction, preserve_topology=True + surface_reduction_rate, preserve_topology=True ) self._annotate(export_surface, group, label_names, label_ids, color) surfaces[group] = export_surface @@ -328,9 +328,9 @@ def process( ) if label_surface is None: continue - if surface_target_reduction > 0.0: + if surface_reduction_rate > 0.0: label_surface = label_surface.decimate_pro( - surface_target_reduction, preserve_topology=True + surface_reduction_rate, preserve_topology=True ) self._annotate( label_surface, group, [label_name], [label_id], color diff --git a/src/physiotwin4d/workflow_convert_vtk_to_usd.py b/src/physiotwin4d/workflow_convert_vtk_to_usd.py index a920a7f5..c05dc5f9 100644 --- a/src/physiotwin4d/workflow_convert_vtk_to_usd.py +++ b/src/physiotwin4d/workflow_convert_vtk_to_usd.py @@ -12,11 +12,13 @@ from pathlib import Path from typing import Any, Literal, Mapping, Optional, Sequence, Union +import numpy as np import pyvista as pv import vtk from .convert_vtk_to_usd import ConvertVTKToUSD from .physiotwin4d_base import PhysioTwin4DBase +from .segment_anatomy_base import SegmentAnatomyBase from .usd_anatomy_tools import USDAnatomyTools from .usd_tools import USDTools @@ -45,6 +47,8 @@ def __init__( solid_color: tuple[float, float, float] = (0.8, 0.8, 0.8), anatomy_type: Optional[str] = None, object_names: Optional[Sequence[str]] = None, + label_names: Optional[Mapping[int, str]] = None, + segmenter: Optional[SegmentAnatomyBase] = None, colormap_primvar: Optional[str] = None, colormap_name: str = "viridis", colormap_intensity_range: Optional[tuple[float, float]] = None, @@ -91,6 +95,25 @@ def __init__( exactly one name (as written by :class:`WorkflowConvertImageToVTK`), and falls back to ``{usd_project_name}_{index}`` otherwise. + label_names: Mapping of label id → structure name. Each input mesh + is then split on its per-cell ``SegmentationLabelIds`` (or + ``boundary_labels``) array, so every structure becomes its own + prim at ``/World/{usd_project_name}/{group}/{structure}``: + time-varying across frames, or static from a single mesh. This + is the only way structure identity survives a time series; + without it, parts are named by connectivity-component order, + which is positional per frame. None (default) reads the ids off + the meshes themselves when they carry the per-cell array, and + names them from *segmenter*'s taxonomy — so passing a mesh + merged by :meth:`ContourTools.save_combined_surfaces` splits by + structure without further arguments. ``static_merge`` accepts + only one labeled mesh: several would collide on one prim path + per label. + segmenter: Segmenter whose taxonomy groups the labels of + *label_names* by anatomy type. Also selects each structure's + material when appearance == "anatomy", through + :meth:`USDAnatomyTools.enhance_meshes`, which falls back to the + containing group for a structure with no material of its own. colormap_primvar: Primvar name for coloring when appearance == "colormap" (e.g. vtk_point_stress_c0). If None, a candidate is auto-picked when possible. colormap_name: Matplotlib colormap name when appearance == "colormap". @@ -117,6 +140,8 @@ def __init__( self.solid_color = solid_color self.anatomy_type = anatomy_type self.object_names = list(object_names) if object_names is not None else None + self.label_names = dict(label_names) if label_names is not None else None + self.segmenter = segmenter self.colormap_primvar = colormap_primvar self.colormap_name = colormap_name self.colormap_intensity_range = colormap_intensity_range @@ -126,6 +151,85 @@ def __init__( "separate_by_connectivity and separate_by_cell_type cannot both be True" ) + @staticmethod + def _as_pyvista( + mesh: Union[pv.DataSet, vtk.vtkDataSet], + ) -> Optional[pv.DataSet]: + """Return *mesh* as a PyVista dataset, or ``None`` if it is not one.""" + if not isinstance(mesh, pv.DataSet) and isinstance(mesh, vtk.vtkDataSet): + mesh = pv.wrap(mesh) + return mesh if isinstance(mesh, pv.DataSet) else None + + def _resolve_label_names(self) -> Optional[dict[int, str]]: + """Return the label ids to split every input mesh on, or ``None``. + + An explicit ``label_names`` is used as given. Otherwise the meshes are + searched for the per-cell label array that + :meth:`ContourTools.save_combined_surfaces` writes on a merge, and that + contouring a multi-label labelmap leaves behind. That array is + preferred wherever it exists because it survives merging, which the + per-object ``field_data`` naming does not — so a combined surface file + splits back into its structures instead of collapsing onto one prim. + + Ids are named from *segmenter*'s taxonomy first, then from the + ``field_data`` of any input holding exactly one structure, and finally + as ``label_{id}``. Ids that no source can name are not worth splitting + on, so a set where none resolve falls back to per-object naming. + + Returns: + The id → name mapping, or ``None`` when no mesh carries the array + or none of its ids can be named, in which case prims are named per + object as before. + """ + if self.label_names is not None: + return self.label_names + + label_ids: set[int] = set() + field_names: dict[int, str] = {} + for mesh in self.input_meshes: + pv_mesh = self._as_pyvista(mesh) + if pv_mesh is None: + continue + for array_name in ("SegmentationLabelIds", "boundary_labels"): + if array_name in pv_mesh.cell_data: + label_ids.update( + int(value) for value in np.unique(pv_mesh.cell_data[array_name]) + ) + break + ids = pv_mesh.field_data.get("SegmentationLabelIds") + names = pv_mesh.field_data.get("SegmentationLabelNames") + if ids is not None and names is not None and len(ids) == len(names) == 1: + field_names[int(ids[0])] = str(names[0]) + + # 0 tags the cells save_combined_surfaces could not attribute to one + # structure, so it names nothing. + label_ids.discard(0) + if not label_ids: + return None + + taxonomy_names = ( + self.segmenter.taxonomy.all_labels() if self.segmenter is not None else {} + ) + named = { + label_id: taxonomy_names.get(label_id) or field_names.get(label_id) + for label_id in sorted(label_ids) + } + if not any(named.values()): + # Ids nobody can name would only produce "label_37" prims, which + # carry less meaning than the object names they would replace. + self.log_debug( + "Per-cell labels %s match no name; naming per object instead", + sorted(label_ids), + ) + return None + resolved = { + label_id: name or f"label_{label_id}" for label_id, name in named.items() + } + self.log_info( + "Splitting on the per-cell label array: %s", ", ".join(resolved.values()) + ) + return resolved + def _read_object_annotations(self) -> list[tuple[Optional[str], Optional[str]]]: """Return ``(structure name, anatomy group)`` per input mesh. @@ -136,11 +240,11 @@ def _read_object_annotations(self) -> list[tuple[Optional[str], Optional[str]]]: """ annotations: list[tuple[Optional[str], Optional[str]]] = [] for mesh in self.input_meshes: - if not isinstance(mesh, pv.DataSet) and isinstance(mesh, vtk.vtkDataSet): - mesh = pv.wrap(mesh) - if not isinstance(mesh, pv.DataSet): + pv_mesh = self._as_pyvista(mesh) + if pv_mesh is None: annotations.append((None, None)) continue + mesh = pv_mesh label_names = mesh.field_data.get("SegmentationLabelNames") groups = mesh.field_data.get("AnatomyGroup") name = ( @@ -212,11 +316,16 @@ def process(self) -> dict[str, Any]: else "none" ) + # Per-cell labels, when the meshes carry them, name the prims instead: + # one per structure, in both the static and the time-series layout. + label_names = self._resolve_label_names() + name_objects = self.static_merge and label_names is None + # Object names only name prims in the static-merge layout; a time # series writes one prim per part across all frames instead. annotations = self._read_object_annotations() object_names = None - if self.static_merge: + if name_objects: object_names = self.object_names if object_names is None and any(name for name, _ in annotations): object_names = [ @@ -231,7 +340,7 @@ def process(self) -> dict[str, Any]: # Keyed by the prim names ConvertVTKToUSD will actually emit, which fall # back to "_" when no object_names were derived. object_groups: dict[str, str] = {} - if self.static_merge: + if name_objects: group_keys = object_names or [ f"{self.usd_project_name}_{index}" for index in range(len(annotations)) ] @@ -242,6 +351,8 @@ def process(self) -> dict[str, Any]: converter = ConvertVTKToUSD( data_basename=self.usd_project_name, input_polydata=self.input_meshes, + mask_ids=label_names, + segmenter=self.segmenter, convert_to_surface=self.extract_surface, separate_by=separate_by, frames_per_second=self.frames_per_second, @@ -281,6 +392,19 @@ def process(self) -> dict[str, Any]: bind_vertex_color_material=True, ) + elif ( + self.appearance == "anatomy" + and label_names is not None + and self.segmenter is not None + ): + # The label layout names each prim after its structure, and the + # segmenter's taxonomy supplies the group to fall back on when the + # structure has no material of its own. + USDAnatomyTools(stage, log_level=self.log_level).enhance_meshes( + self.segmenter + ) + stage.Save() + elif self.appearance == "anatomy": anatomy_tools = USDAnatomyTools(stage, log_level=self.log_level) for mesh_path in mesh_paths: diff --git a/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py b/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py index ea97d1ef..614da426 100644 --- a/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py +++ b/src/physiotwin4d/workflow_fit_statistical_model_to_patient.py @@ -80,6 +80,9 @@ class WorkflowFitStatisticalModelToPatient(PhysioTwin4DBase): patient_labelmap (itk.Image): Multi-label labelmap for patient model patient_mask (itk.Image): Binary mask for patient registration region mask_dilation_mm (float): Dilation for binary mask generation + distancemap_squared_max (Optional[float]): Saturation radius of the + labelmap-to-labelmap distance maps, in squared millimeters. None + means derive it from mask_dilation_mm as (1.25 * mask_dilation_mm)**2 transform_tools (TransformTools): Transform utilities registrar_ICON (RegisterImagesICON): ICON registration instance registrar_Greedy (RegisterImagesGreedy): Greedy registration instance diff --git a/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py b/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py index b0c28036..f354bc2c 100644 --- a/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py +++ b/src/physiotwin4d/workflow_reconstruct_highres_4d_ct.py @@ -4,7 +4,7 @@ a high-resolution 4D CT time series from lower-resolution time-series images and a single high-resolution reference image. -The workflow uses Greedy+ICON combined registration to: +The workflow uses a caller-supplied RegisterImagesBase backend to: 1. Register each time-series image to the high-resolution reference 2. Apply inverse transforms to reconstruct high-resolution time series 3. Optionally upsample to the reference image resolution @@ -16,7 +16,7 @@ Key Features: - Sequential time-series registration using RegisterTimeSeriesImages - - Combined Greedy+ICON registration for optimal results + - Any RegisterImagesBase backend, defaulting to RegisterImagesGreedy - Bidirectional registration from reference frame - Optional temporal smoothing with prior transforms - High-resolution reconstruction with optional upsampling @@ -38,7 +38,7 @@ class WorkflowReconstructHighres4DCT(PhysioTwin4DBase): This class implements a workflow for reconstructing high-resolution dynamic CT images by registering low-resolution time-series images to a high-resolution - reference image using combined Greedy+ICON registration. + reference image using a caller-supplied registration backend. **Registration Pipeline:** 1. **Time Series Registration**: Register each time-series image to the @@ -109,7 +109,7 @@ def __init__( an identity transform for that frame. Default: True registration_method (Optional[RegisterImagesBase]): Registration backend instance. Defaults to a new - :class:`RegisterImagesGreedyICON` when None. + :class:`RegisterImagesGreedy` when None. log_level: Logging level (logging.DEBUG, logging.INFO, etc.). Default: logging.INFO diff --git a/tests/test_contour_mesh_extraction.py b/tests/test_contour_mesh_extraction.py new file mode 100644 index 00000000..862eb083 --- /dev/null +++ b/tests/test_contour_mesh_extraction.py @@ -0,0 +1,401 @@ +"""Tests for ContourTools' per-label surface and tetrahedral mesh extraction. + +These use ``data/test/slicer_heart_small``, whose direction matrix is +right-handed, plus a synthetic left-handed image for the cases that only a +negative-determinant direction can exercise. +""" + +from __future__ import annotations + +from pathlib import Path + +import itk +import numpy as np +import pytest +import pyvista as pv + +from physiotwin4d.contour_tools import ContourTools + +#: A label that reaches the volume border, so its surface only closes if the +#: mask is padded first. +BORDER_LABEL = 5 + + +@pytest.fixture(scope="module") +def heart_labelmap() -> itk.Image: + """Read one small labelmap.""" + path = ( + Path(__file__).resolve().parents[1] + / "data" + / "test" + / "slicer_heart_small" + / "slice_000_labelmap.mha" + ) + if not path.exists(): + pytest.skip(f"Test labelmap not available: {path}") + return itk.imread(str(path)) + + +def _label_mask(labelmap: itk.Image, label_id: int) -> itk.Image: + """Isolate one label of *labelmap* as a binary mask.""" + arr = itk.GetArrayViewFromImage(labelmap) + mask = itk.GetImageFromArray((arr == label_id).astype(np.uint8)) + mask.CopyInformation(labelmap) + return mask + + +def _left_handed_box() -> itk.Image: + """Build a solid box in an image whose direction determinant is negative.""" + arr = np.zeros((6, 6, 6), dtype=np.uint8) + arr[1:5, 1:5, 1:5] = 1 + mask = itk.GetImageFromArray(arr) + mask.SetSpacing([1.0, 2.0, 3.0]) + mask.SetDirection(itk.matrix_from_array(np.diag([1.0, 1.0, -1.0]))) + return mask + + +def _touching_boxes() -> itk.Image: + """Build two boxes sharing a wall, sampled coarsely along the slice axis. + + The 4 mm slice pitch against the 1 mm in-plane pitch is what makes an + independently contoured surface terrace, and the negative-determinant + direction is the LPS case that inverts VTK's face winding. + """ + arr = np.zeros((14, 48, 60), dtype=np.uint16) + arr[3:11, 8:40, 6:30] = 1 + arr[3:11, 8:40, 30:54] = 2 + labelmap = itk.GetImageFromArray(arr) + labelmap.SetSpacing([1.0, 1.0, 4.0]) + labelmap.SetDirection(itk.matrix_from_array(np.diag([1.0, -1.0, 1.0]))) + return labelmap + + +def _point_keys(surface: pv.PolyData) -> set[tuple[float, ...]]: + """Return *surface*'s points as hashable, rounded coordinate triples.""" + return {tuple(point) for point in np.round(np.asarray(surface.points), 6).tolist()} + + +def _roughness(surface: pv.PolyData) -> float: + """Return how far a point sits from its neighbors' centroid, per edge length. + + Dividing by the mean edge length makes the measure comparable between + surfaces triangulated at different resolutions. + """ + mesh = surface.triangulate() + faces = mesh.faces.reshape(-1, 4)[:, 1:] + points = np.asarray(mesh.points) + total = np.zeros_like(points) + count = np.zeros(len(points)) + for first, second in ((0, 1), (1, 2), (2, 0)): + for source, target in ((first, second), (second, first)): + np.add.at(total, faces[:, source], points[faces[:, target]]) + np.add.at(count, faces[:, source], 1) + offset = total / np.maximum(count, 1)[:, np.newaxis] - points + edge_length = np.linalg.norm(points[faces[:, 0]] - points[faces[:, 1]], axis=1) + return float(np.linalg.norm(offset, axis=1).mean() / edge_length.mean()) + + +def _label_pairs(contours: pv.PolyData) -> set[tuple[int, ...]]: + """Return the set of label pairs the contour's cells separate.""" + pairs = np.asarray(contours.cell_data["boundary_labels"]) + return {tuple(row) for row in np.unique(pairs.reshape(len(pairs), -1), axis=0)} + + +class TestExtractContours: + """The multi-label contour must be smooth and keep its labels in contact.""" + + def test_smoothing_takes_the_voxel_blocks_out( + self, contour_tools: ContourTools + ) -> None: + """Taubin smoothing leaves the surface markedly less faceted. + + These boxes are flat over most of their area, where there is nothing to + take out, so the margin here is smaller than on anatomy: the same + measure halves on a segmented chest CT. + """ + blocky = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=0) + smooth = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=20) + + assert _roughness(smooth) < 0.75 * _roughness(blocky) + + def test_smoothing_keeps_every_cell_and_its_labels( + self, contour_tools: ContourTools + ) -> None: + """Smoothing only moves points, so the label of each cell survives. + + Downstream splits a contour into one prim per structure on + ``boundary_labels``, so a lost pair is a lost structure. + """ + blocky = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=0) + smooth = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=20) + + assert smooth.n_cells == blocky.n_cells + assert smooth.n_points == blocky.n_points + assert _label_pairs(smooth) == _label_pairs(blocky) + + def test_labels_stay_in_contact(self, contour_tools: ContourTools) -> None: + """The wall between the labels is not torn open by smoothing. + + Regression guard for ``non_manifold_smoothing``: the surface net shares + its points between neighboring labels, so smoothing the mesh as one + keeps them together, and the edges where three labels meet must move + with the rest instead of being pinned. + """ + blocky = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=0) + smooth = contour_tools.extract_contours(_touching_boxes(), taubin_iterations=20) + + def open_edges(contours: pv.PolyData) -> int: + return int( + contours.extract_feature_edges( + boundary_edges=True, + feature_edges=False, + manifold_edges=False, + non_manifold_edges=False, + ).n_cells + ) + + assert open_edges(smooth) == open_edges(blocky) + + def test_anisotropic_labelmap_is_contoured_isotropically( + self, contour_tools: ContourTools + ) -> None: + """Boundaries land between the slices rather than terracing at them. + + Contoured on the labelmap's own 4 mm grid, every vertex of a face + parallel to the slices lies on a slice plane; the isotropic resample is + what lets one land in between. + """ + labelmap = _touching_boxes() + origin_z = float(np.asarray(labelmap.GetOrigin())[2]) + + contours = contour_tools.extract_contours(labelmap) + + offsets = (np.asarray(contours.points)[:, 2] - origin_z) / 4.0 + assert np.any(np.abs(offsets - np.round(offsets)) > 0.1) + + +class TestExtractLabelSurfaces: + """Labels extracted together must stay closed and stay in contact.""" + + def test_surfaces_are_watertight_and_outward( + self, contour_tools: ContourTools + ) -> None: + """Every label closes and encloses a positive volume.""" + surfaces = contour_tools.extract_label_surfaces(_touching_boxes()) + + assert sorted(surfaces) == [1, 2] + for label_id, surface in surfaces.items(): + assert contour_tools.is_watertight(surface), ( + f"label {label_id} must be closed" + ) + assert surface.volume > 0.0, f"label {label_id} must face outward" + + def test_neighbors_share_the_wall_between_them( + self, contour_tools: ContourTools + ) -> None: + """The two surfaces meet on identical vertices, smoothing included. + + Regression guard for the merge that welds the wall and for smoothing + the welded mesh as one: smoothing the labels separately, or without + ``non_manifold_smoothing``, pulls the two copies of the wall apart. + """ + surfaces = contour_tools.extract_label_surfaces( + _touching_boxes(), smoothing_iterations=20 + ) + + shared = _point_keys(surfaces[1]) & _point_keys(surfaces[2]) + assert len(shared) > 100, "the shared wall must be meshed identically" + distances = ( + surfaces[1].copy().compute_implicit_distance(surfaces[2]) + ).point_data["implicit_distance"] + assert np.count_nonzero(np.abs(distances) < 1e-9) == len(shared) + + def test_volume_matches_the_labelmap(self, contour_tools: ContourTools) -> None: + """Neither label is thinned or fattened by the smoothing. + + Not to the voxel exactly: the distance map is zero at the outermost + labeled voxel's center rather than half a voxel beyond it, so the + surface sits just inside the block of voxels it came from. + """ + labelmap = _touching_boxes() + labels = itk.GetArrayViewFromImage(labelmap) + voxel_volume = float(np.prod(np.asarray(labelmap.GetSpacing()))) + + surfaces = contour_tools.extract_label_surfaces(labelmap) + + for label_id, surface in surfaces.items(): + voxels = int(np.count_nonzero(labels == label_id)) + assert surface.volume == pytest.approx(voxels * voxel_volume, rel=0.15) + + +class TestExtractWatertightSurface: + """The per-label surface must be closed and outward-oriented.""" + + def test_border_label_surface_is_watertight( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """A label touching the volume border still closes. + + Regression guard for the one-voxel background pad: without it the + isosurface is cut open where the structure reaches the image edge. + """ + surface = contour_tools.extract_watertight_surface( + _label_mask(heart_labelmap, BORDER_LABEL) + ) + + assert surface.n_cells > 0 + assert contour_tools.is_watertight(surface), ( + "every edge must be shared by exactly two faces" + ) + + def test_normals_point_outward(self, contour_tools: ContourTools) -> None: + """A left-handed direction still yields a positive enclosed volume. + + Regression guard for ``auto_orient_normals``: VTK winds faces for a + right-handed direction matrix, so LPS images with a negative-determinant + direction otherwise come out with inward normals. + """ + surface = contour_tools.extract_watertight_surface(_left_handed_box()) + + volume = surface.triangulate().compute_cell_sizes(volume=True) + assert float(np.sum(volume["Volume"])) == pytest.approx(0.0, abs=1e-9), ( + "cell volumes of a surface sum to zero; enclosed volume is checked below" + ) + assert surface.volume > 0.0, "enclosed volume must be positive" + + def test_decimation_reduces_triangle_count( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """surface_reduction_rate removes roughly that fraction of triangles.""" + mask = _label_mask(heart_labelmap, BORDER_LABEL) + full = contour_tools.extract_watertight_surface(mask) + reduced = contour_tools.extract_watertight_surface( + mask, surface_reduction_rate=0.5 + ) + + assert reduced.n_cells == pytest.approx(full.n_cells / 2, rel=0.1) + + +class TestExtractTetrahedra: + """The per-label volume mesh must be tetrahedral and positively oriented.""" + + def test_cells_are_tetrahedra( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """Every cell is a tetrahedron, six per labeled voxel.""" + mask = _label_mask(heart_labelmap, BORDER_LABEL) + mesh = contour_tools.extract_tetrahedra(mask) + + assert set(mesh.celltypes) == {pv.CellType.TETRA} + voxels = int(np.count_nonzero(itk.GetArrayViewFromImage(mask))) + assert mesh.n_cells == 6 * voxels + + def test_left_handed_direction_yields_positive_volumes( + self, contour_tools: ContourTools + ) -> None: + """No tetrahedron is inverted when the direction determinant is negative. + + Regression guard for the hexahedron winding flip: triangulating the raw + voxel cells of a left-handed image inverts every tetrahedron. + """ + mesh = contour_tools.extract_tetrahedra(_left_handed_box()) + + volumes = mesh.compute_cell_sizes(volume=True)["Volume"] + assert np.all(volumes > 0.0), "no tetrahedron may have negative volume" + # 4 x 4 x 4 voxels of 1 x 2 x 3 mm. + assert float(np.sum(volumes)) == pytest.approx(4 * 4 * 4 * 6.0) + + def test_volume_survives_coarsening( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """Elements twice the voxel keep the volume while shedding cells cubically.""" + mask = _label_mask(heart_labelmap, BORDER_LABEL) + voxel = float(np.min(np.asarray(mask.GetSpacing()))) + full = contour_tools.extract_tetrahedra(mask, element_size_mm=voxel) + coarse = contour_tools.extract_tetrahedra(mask, element_size_mm=2.0 * voxel) + + full_volume = float(np.sum(full.compute_cell_sizes(volume=True)["Volume"])) + coarse_volume = float(np.sum(coarse.compute_cell_sizes(volume=True)["Volume"])) + assert coarse_volume == pytest.approx(full_volume, rel=0.1) + assert coarse.n_cells == pytest.approx(full.n_cells * 0.125, rel=0.3) + + def test_elements_are_isotropic(self, contour_tools: ContourTools) -> None: + """The requested element size is what the mesh is built on. + + Regression guard for the anisotropic default: meshing a mask's own + voxels inherits their slice pitch, which is what made these meshes + blocky along one axis. + """ + mesh = contour_tools.extract_tetrahedra(_left_handed_box(), element_size_mm=1.0) + + volumes = np.asarray(mesh.compute_cell_sizes(volume=True)["Volume"]) + # Six tetrahedra fill a 1 mm cube. + assert float(np.median(volumes)) == pytest.approx(1.0 / 6.0, rel=1e-6) + + +class TestTrimTetrahedraToSurface: + """Relaxing must land the boundary on the surface without wrecking cells.""" + + @staticmethod + def _boundary_gap(mesh: pv.UnstructuredGrid, surface: pv.PolyData) -> float: + """Return the mean distance, in mm, from *mesh*'s boundary to *surface*.""" + boundary = mesh.extract_surface(algorithm="dataset_surface") + distance = np.asarray( + boundary.compute_implicit_distance(surface).point_data["implicit_distance"] + ) + return float(np.abs(distance).mean()) + + def test_boundary_lands_on_the_surface( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """The voxel staircase is relaxed onto the surface, not merely clipped. + + Regression guard for the interior smoothing: projecting the boundary + alone flattens the cells behind it, the quality bound undoes the move, + and the staircase survives. + """ + mask = _label_mask(heart_labelmap, BORDER_LABEL) + surface = contour_tools.extract_label_surfaces(mask)[1] + mesh = contour_tools.extract_tetrahedra(mask) + relaxed = contour_tools.trim_tetrahedra_to_surface(mesh, surface) + + before = self._boundary_gap(mesh, surface) + assert before > 0.0, "the staircase must start off the surface" + assert self._boundary_gap(relaxed, surface) < 0.25 * before + + def test_cells_are_kept_whole_and_well_shaped( + self, contour_tools: ContourTools, heart_labelmap: itk.Image + ) -> None: + """Cells are only dropped, never cut into slivers. + + Regression guard against swapping the crinkle clip for a real + ``clip_surface`` cut: cutting subdivides the boundary tetrahedra and + drives about a tenth of them below a scaled Jacobian of 0.1. + """ + mask = _label_mask(heart_labelmap, BORDER_LABEL) + surface = contour_tools.extract_label_surfaces(mask)[1] + mesh = contour_tools.extract_tetrahedra(mask) + relaxed = contour_tools.trim_tetrahedra_to_surface(mesh, surface) + + assert set(relaxed.celltypes) == {pv.CellType.TETRA} + assert 0 < relaxed.n_cells <= mesh.n_cells + quality = np.asarray( + relaxed.cell_quality(["scaled_jacobian"]).cell_data["scaled_jacobian"] + ) + assert np.min(quality) >= 0.1, "no tetrahedron may be left a sliver" + + def test_anatomy_color_survives(self, contour_tools: ContourTools) -> None: + """Cell and field data attached at extraction are carried through.""" + mask = _left_handed_box() + surface = contour_tools.extract_label_surfaces(mask)[1] + mesh = contour_tools.extract_tetrahedra(mask, anatomy_names=["heart"]) + relaxed = contour_tools.trim_tetrahedra_to_surface(mesh, surface) + + assert "Color" in relaxed.cell_data + assert np.array_equal( + relaxed.field_data["AnatomyColor"], mesh.field_data["AnatomyColor"] + ) + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_convert_vtk_to_usd.py b/tests/test_convert_vtk_to_usd.py index 2ab2c68d..27d4e30f 100644 --- a/tests/test_convert_vtk_to_usd.py +++ b/tests/test_convert_vtk_to_usd.py @@ -355,16 +355,21 @@ def test_time_varying_topology( assert stage is not None, "USD stage not created" assert output_file.exists(), "USD file not created" - # Check for time-varying meshes (separate mesh prims) - parent_path = "/HeartVarying/default" - parent_prim = stage.GetPrimAtPath(parent_path) - - # Should have child meshes for each time step - children = parent_prim.GetChildren() if parent_prim.IsValid() else [] + # Each frame carries its own topology, so neither one is indexed by the + # other's faces. + mesh = UsdGeom.Mesh(stage.GetPrimAtPath("/World/HeartVarying/Mesh")) + assert mesh.GetPrim().IsValid(), "Mesh prim not found" + indices_attr = mesh.GetFaceVertexIndicesAttr() + assert indices_attr.GetTimeSamples() == [0.0, 1.0] + for time_code in (0.0, 1.0): + points = mesh.GetPointsAttr().Get(time_code) + indices = indices_attr.Get(time_code) + assert max(indices) < len(points), ( + f"Face indices at t={time_code} address points that frame " + f"does not have ({max(indices)} >= {len(points)})" + ) print("Time-varying topology handled") - print(f" Parent prim: {parent_path}") - print(f" Child prims: {len(children)}") print(f" Output: {output_file}") def test_batch_conversion( @@ -477,6 +482,48 @@ def test_static_merge_prim_names_use_data_basename(self, tmp_path: Path) -> None f"{prim_path} should have no time samples but got {samples}" ) + # ------------------------------------------------------------------ + # Frames that share no triangulation must each author their own topology + # ------------------------------------------------------------------ + + def test_varying_topology_time_samples_faces(self, tmp_path: Path) -> None: + """Independently built frames must not be indexed by frame 0's faces.""" + frames = [ + pv.Sphere(theta_resolution=8 + index, phi_resolution=8) + for index in range(3) + ] + converter = ConvertVTKToUSD(data_basename="P", input_polydata=frames) + stage = converter.convert(str(tmp_path / "out.usd")) + + mesh = UsdGeom.Mesh(stage.GetPrimAtPath("/World/P/Mesh")) + counts_attr = mesh.GetFaceVertexCountsAttr() + indices_attr = mesh.GetFaceVertexIndicesAttr() + assert counts_attr.GetTimeSamples() == [0.0, 1.0, 2.0] + assert indices_attr.GetTimeSamples() == [0.0, 1.0, 2.0] + + for time_code in (0.0, 1.0, 2.0): + points = mesh.GetPointsAttr().Get(time_code) + indices = indices_attr.Get(time_code) + counts = counts_attr.Get(time_code) + assert max(indices) < len(points), ( + f"Face indices at t={time_code} address points that frame " + f"does not have ({max(indices)} >= {len(points)})" + ) + assert sum(counts) == len(indices) + + def test_constant_topology_leaves_faces_untimed(self, tmp_path: Path) -> None: + """A deformed series keeps one topology, so viewers can interpolate.""" + base = _make_poly() + moved = base.copy() + moved.points[:, 2] += 1.0 + converter = ConvertVTKToUSD(data_basename="P", input_polydata=[base, moved]) + stage = converter.convert(str(tmp_path / "out.usd")) + + mesh = UsdGeom.Mesh(stage.GetPrimAtPath("/World/P/Mesh")) + assert mesh.GetFaceVertexCountsAttr().GetTimeSamples() == [] + assert mesh.GetFaceVertexIndicesAttr().GetTimeSamples() == [] + assert mesh.GetPointsAttr().GetTimeSamples() == [0.0, 1.0] + def test_mask_ids_split_on_segmentation_label_ids(self, tmp_path: Path) -> None: """A merged surface file splits on the array save_combined_surfaces writes.""" mesh = _make_poly() diff --git a/tests/test_image_tools.py b/tests/test_image_tools.py index 13af8595..006da590 100644 --- a/tests/test_image_tools.py +++ b/tests/test_image_tools.py @@ -484,5 +484,59 @@ def test_keep_largest_connected_component_no_foreground_returns_empty( assert result_arr.sum() == 0 +class TestResampleImageByScale: + """Unit tests for ImageTools.resample_image_by_scale.""" + + @pytest.fixture + def image_tools(self) -> ImageTools: + return ImageTools() + + @pytest.mark.parametrize("scale", [0.5, 0.25, 2.0]) + def test_physical_extent_is_preserved( + self, image_tools: ImageTools, scale: float + ) -> None: + """Voxel count scales while the physical extent stays put.""" + itk_image = _make_synthetic_itk_image((8, 6, 4)) + itk_image.SetSpacing([1.0, 2.0, 3.0]) + + out = image_tools.resample_image_by_scale(itk_image, scale) + + size = np.asarray(itk_image.GetLargestPossibleRegion().GetSize()) + new_size = np.asarray(out.GetLargestPossibleRegion().GetSize()) + assert np.array_equal(new_size, np.ceil(size * scale).astype(int)) + assert np.allclose( + new_size * np.asarray(out.GetSpacing()), + size * np.asarray(itk_image.GetSpacing()), + ), "resampling must not change the physical extent" + + def test_direction_is_unchanged(self, image_tools: ImageTools) -> None: + """A left-handed direction survives resampling untouched.""" + direction = np.diag([1.0, 1.0, -1.0]) + itk_image = _make_synthetic_itk_image((8, 6, 4), direction=direction) + + out = image_tools.resample_image_by_scale(itk_image, 0.5) + + assert np.allclose(itk.array_from_matrix(out.GetDirection()), direction) + + def test_nearest_neighbor_keeps_input_values(self, image_tools: ImageTools) -> None: + """Nearest-neighbor resampling introduces no new intensities.""" + arr = np.zeros((4, 4, 4), dtype=np.float32) + arr[1:3, 1:3, 1:3] = 7.0 + itk_image = _make_synthetic_itk_image((4, 4, 4), arr=arr) + + out = image_tools.resample_image_by_scale(itk_image, 0.5, interpolate=False) + + assert set(np.unique(itk.array_from_image(out))) <= {0.0, 7.0} + + @pytest.mark.parametrize("scale", [0.0, -1.0]) + def test_non_positive_scale_raises( + self, image_tools: ImageTools, scale: float + ) -> None: + """A scale of zero or less is rejected.""" + itk_image = _make_synthetic_itk_image((4, 4, 4)) + with pytest.raises(ValueError, match="scale must be positive"): + image_tools.resample_image_by_scale(itk_image, scale) + + if __name__ == "__main__": pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_register_images_base.py b/tests/test_register_images_base.py new file mode 100644 index 00000000..cf91bf20 --- /dev/null +++ b/tests/test_register_images_base.py @@ -0,0 +1,116 @@ +"""Tests for the background value RegisterImagesBase writes off-grid. + +The images here are synthetic: a moving image shifted off the fixed grid, so +every resampling leaves a region the moving image does not cover. +""" + +import itk +import numpy as np + +from physiotwin4d.register_images_base import RegisterImagesBase + +_SIZE = 10 +_SHIFT_MM = 5.0 + + +def _uniform_image(value: float, dtype: type) -> itk.Image: + """Return a cubic image whose voxels all hold ``value``.""" + return itk.image_from_array(np.full((_SIZE, _SIZE, _SIZE), value, dtype=dtype)) + + +def _shift_transform() -> itk.Transform: + """Return a translation large enough to push half the grid off the moving image.""" + transform = itk.TranslationTransform[itk.D, 3].New() + transform.SetOffset([_SHIFT_MM, 0.0, 0.0]) + return transform + + +def _off_grid_voxel(image: itk.Image) -> float: + """Return the voxel that the shift maps outside the moving image.""" + return float(itk.GetArrayViewFromImage(image)[0, 0, _SIZE - 1]) + + +def _on_grid_voxel(image: itk.Image) -> float: + """Return a voxel the shift keeps inside the moving image.""" + return float(itk.GetArrayViewFromImage(image)[0, 0, 0]) + + +def test_prewarp_background_value_uses_override() -> None: + """An explicit override wins over the modality default.""" + registrar = RegisterImagesBase() + registrar.set_modality("ct") + registrar.set_prewarp_background_value(42.0) + + assert ( + registrar._prewarp_background_value(_uniform_image(100.0, np.float32)) == 42.0 + ) + + +def test_prewarp_background_value_ct_is_air() -> None: + """CT falls back to -1000 HU rather than the image's own minimum.""" + registrar = RegisterImagesBase() + registrar.set_modality("ct") + + assert ( + registrar._prewarp_background_value(_uniform_image(100.0, np.float32)) + == -1000.0 + ) + + +def test_prewarp_background_value_non_ct_uses_image_minimum() -> None: + """Other modalities have no fixed air value, so the image's minimum is used.""" + registrar = RegisterImagesBase() + registrar.set_modality("mri") + moving_image = _uniform_image(100.0, np.float32) + itk.GetArrayViewFromImage(moving_image)[0, 0, 0] = 7.0 + + assert registrar._prewarp_background_value(moving_image) == 7.0 + + +def test_prewarp_moving_fills_image_with_modality_background() -> None: + """The pre-warped intensity image is filled with air, not 0 HU.""" + registrar = RegisterImagesBase() + registrar.set_modality("ct") + registrar.set_fixed_image(_uniform_image(0.0, np.float32)) + + warped_image, warped_mask, warped_labelmap = registrar._prewarp_moving( + _shift_transform(), _uniform_image(100.0, np.float32), None, None + ) + + assert _on_grid_voxel(warped_image) == 100.0 + assert _off_grid_voxel(warped_image) == -1000.0 + assert warped_mask is None + assert warped_labelmap is None + + +def test_prewarp_moving_fills_mask_and_labelmap_with_zero() -> None: + """Masks and labelmaps keep their discrete values and are filled with 0.""" + registrar = RegisterImagesBase() + registrar.set_modality("ct") + registrar.set_fixed_image(_uniform_image(0.0, np.float32)) + + _, warped_mask, warped_labelmap = registrar._prewarp_moving( + _shift_transform(), + _uniform_image(100.0, np.float32), + _uniform_image(1, np.uint8), + _uniform_image(3, np.uint8), + ) + + assert _on_grid_voxel(warped_mask) == 1 + assert _off_grid_voxel(warped_mask) == 0 + assert _on_grid_voxel(warped_labelmap) == 3 + assert _off_grid_voxel(warped_labelmap) == 0 + + +def test_get_registered_image_fills_with_modality_background() -> None: + """The registered CT is filled with air where the moving image ends.""" + registrar = RegisterImagesBase() + registrar.set_modality("ct") + registrar.set_fixed_image(_uniform_image(0.0, np.float32)) + registrar.moving_image = _uniform_image(100.0, np.float32) + registrar.forward_transform = _shift_transform() + + registered = registrar.get_registered_image() + + assert _on_grid_voxel(registered) == 100.0 + assert _off_grid_voxel(registered) == -1000.0 diff --git a/tests/test_tutorials.py b/tests/test_tutorials.py index 9f15ca1b..7e1da27a 100644 --- a/tests/test_tutorials.py +++ b/tests/test_tutorials.py @@ -345,9 +345,11 @@ def test_run(self, test_directories: dict[str, Path]) -> None: assert (model_dir / "mgn_stage_model.pt").exists(), "Checkpoint should exist" assert results["cases"], "At least one held-out case should be evaluated" + # The model goes to the shared weights directory; the manifests, the + # evaluation and the screenshots stay under the tutorial's output. tt = TestTools( class_name=self._class_name, - results_dir=model_dir, + results_dir=_REPO_ROOT / "tutorials" / "output" / "tutorial_09_lung_mgn", baselines_dir=test_directories["baselines"] / self._class_name, ) _compare_screenshots(results["screenshots"], tt) @@ -364,7 +366,10 @@ class TestTutorial10LungInferPhysicsNeMoMGN: def test_run(self, test_directories: dict[str, Path]) -> None: _require_physicsnemo_and_tutorial_08() - model_dir = _REPO_ROOT / "tutorials" / "output" / "tutorial_09_lung_mgn" + # ParametersLungCTDirLab.mgn_weights_dir, where Tutorial 9 trains to. + model_dir = ( + _REPO_ROOT / "tutorials" / "network_weights" / "physicsnemo_mgn_lung_motion" + ) if not (model_dir / "mgn_stage_model.pt").exists(): _run_tutorial_script("tutorial_09_lung_train_physicsnemo_mgn.py") assert (model_dir / "mgn_stage_model.pt").exists(), ( @@ -377,7 +382,9 @@ def test_run(self, test_directories: dict[str, Path]) -> None: ) assert Path(results["usd_file"]).exists(), "USD file should exist" - out_dir = model_dir / "tutorial_10_lung_mgn" / "Case1Pack" + out_dir = ( + _REPO_ROOT / "tutorials" / "output" / "tutorial_10_lung_mgn" / "Case1Pack" + ) tt = TestTools( class_name=self._class_name, results_dir=out_dir, diff --git a/tests/test_workflow_convert_vtk_to_usd.py b/tests/test_workflow_convert_vtk_to_usd.py index 1fc6ce2a..3128c08e 100644 --- a/tests/test_workflow_convert_vtk_to_usd.py +++ b/tests/test_workflow_convert_vtk_to_usd.py @@ -4,13 +4,14 @@ """ from pathlib import Path -from typing import Optional +from typing import Optional, cast import numpy as np +import pytest import pyvista as pv -from pxr import Usd, UsdShade +from pxr import Usd, UsdGeom, UsdShade -from physiotwin4d import WorkflowConvertVTKToUSD +from physiotwin4d import SegmentHeartSimpleware, WorkflowConvertVTKToUSD def _labeled_sphere( @@ -26,6 +27,28 @@ def _labeled_sphere( return surface +def _two_structure_frame(theta_resolution: int) -> pv.PolyData: + """Return one frame holding two structures tagged per cell by label id. + + *theta_resolution* varies the triangulation, so a series built from several + frames shares neither point count nor face count -- the case of surfaces + contoured independently per frame rather than propagated by registration. + """ + parts = [] + for label_id, center in ((1, (0.0, 0.0, 0.0)), (5, (5.0, 0.0, 0.0))): + part = pv.Sphere( + radius=1.0, + center=center, + theta_resolution=theta_resolution, + phi_resolution=8, + ) + part.cell_data["SegmentationLabelIds"] = np.full( + part.n_cells, label_id, dtype=np.int32 + ) + parts.append(part) + return cast(pv.PolyData, pv.merge(parts, merge_points=False)) + + def _bound_material_path(stage: Usd.Stage, mesh_path: str) -> str: prim = stage.GetPrimAtPath(mesh_path) assert prim.IsValid(), f"Missing prim: {mesh_path}" @@ -167,3 +190,125 @@ def test_unlabeled_meshes_keep_positional_names(self, tmp_path: Path) -> None: stage = Usd.Stage.Open(result["usd_file"]) assert stage.GetPrimAtPath("/World/scan/scan_0_object1").IsValid() assert stage.GetPrimAtPath("/World/scan/scan_1_object1").IsValid() + + +class TestLabelTimeSeries: + """label_names keeps structure identity across the frames of a series.""" + + def test_label_names_produce_per_structure_animated_prims( + self, tmp_path: Path + ) -> None: + """Each label becomes one prim, grouped and painted by the taxonomy.""" + frames = [_two_structure_frame(resolution) for resolution in (8, 9, 10)] + segmenter = SegmentHeartSimpleware() + + workflow = WorkflowConvertVTKToUSD( + input_meshes=frames, + usd_project_name="pm0001", + output_directory=tmp_path, + appearance="anatomy", + label_names={1: "left_ventricle", 5: "myocardium"}, + segmenter=segmenter, + ) + result = workflow.process() + + stage = Usd.Stage.Open(result["usd_file"]) + ventricle_path = "/World/pm0001/heart/left_ventricle" + myocardium_path = "/World/pm0001/heart/myocardium" + + # "left_ventricle" matches no material of its own, so it lands on the + # heart group's; "myocardium" has one. + assert _bound_material_path(stage, ventricle_path).endswith("OmniSurface_Heart") + assert _bound_material_path(stage, myocardium_path).endswith( + "OmniSurface_Myocardium" + ) + + for mesh_path in (ventricle_path, myocardium_path): + mesh = UsdGeom.Mesh(stage.GetPrimAtPath(mesh_path)) + assert mesh.GetPointsAttr().GetTimeSamples() == [0.0, 1.0, 2.0] + + def test_per_cell_labels_are_found_without_label_names( + self, tmp_path: Path + ) -> None: + """The ids on the meshes are enough; the taxonomy supplies the names.""" + workflow = WorkflowConvertVTKToUSD( + input_meshes=[_two_structure_frame(8), _two_structure_frame(9)], + usd_project_name="pm0002", + output_directory=tmp_path, + appearance="anatomy", + segmenter=SegmentHeartSimpleware(), + ) + result = workflow.process() + + stage = Usd.Stage.Open(result["usd_file"]) + assert stage.GetPrimAtPath("/World/pm0002/heart/left_ventricle").IsValid() + assert stage.GetPrimAtPath("/World/pm0002/heart/myocardium").IsValid() + + def test_static_merge_splits_a_labeled_mesh_by_structure( + self, tmp_path: Path + ) -> None: + """One merged surface must not collapse onto a single static prim.""" + workflow = WorkflowConvertVTKToUSD( + input_meshes=[_two_structure_frame(8)], + usd_project_name="patient", + output_directory=tmp_path, + appearance="anatomy", + segmenter=SegmentHeartSimpleware(), + static_merge=True, + ) + result = workflow.process() + + stage = Usd.Stage.Open(result["usd_file"]) + assert stage.GetPrimAtPath("/World/patient/heart/left_ventricle").IsValid() + assert stage.GetPrimAtPath("/World/patient/heart/myocardium").IsValid() + assert not stage.GetPrimAtPath("/World/patient/patient_0_object1").IsValid(), ( + "Per-object naming used despite the per-cell label array" + ) + + def test_unlabeled_meshes_still_name_per_object(self, tmp_path: Path) -> None: + """Without the per-cell array, naming falls back to the old path.""" + workflow = WorkflowConvertVTKToUSD( + input_meshes=[_labeled_sphere((0.0, 0.0, 0.0), "highres_myocardium")], + usd_project_name="heart", + output_directory=tmp_path, + appearance="anatomy", + segmenter=SegmentHeartSimpleware(), + static_merge=True, + ) + result = workflow.process() + + stage = Usd.Stage.Open(result["usd_file"]) + material = _bound_material_path( + stage, "/World/heart/highres_myocardium_object1" + ) + assert material.endswith("OmniSurface_Myocardium") + + def test_nameless_labels_fall_back_to_per_object(self, tmp_path: Path) -> None: + """Without a segmenter the ids name nothing, so splitting on them would + replace "scan_0" with "label_1" and gain nothing.""" + workflow = WorkflowConvertVTKToUSD( + input_meshes=[_two_structure_frame(8)], + usd_project_name="scan", + output_directory=tmp_path, + appearance="anatomy", + anatomy_type="heart", + static_merge=True, + ) + result = workflow.process() + + stage = Usd.Stage.Open(result["usd_file"]) + assert stage.GetPrimAtPath("/World/scan/scan_0_object1").IsValid() + assert not stage.GetPrimAtPath("/World/scan/heart/label_1").IsValid() + + def test_several_labeled_static_meshes_raise(self, tmp_path: Path) -> None: + """Two static objects holding the same labels would collide on one path.""" + workflow = WorkflowConvertVTKToUSD( + input_meshes=[_two_structure_frame(8), _two_structure_frame(9)], + usd_project_name="pm0003", + output_directory=tmp_path, + appearance="anatomy", + segmenter=SegmentHeartSimpleware(), + static_merge=True, + ) + with pytest.raises(ValueError, match="static_merge with mask_ids"): + workflow.process() diff --git a/tutorials/README.md b/tutorials/README.md index 71cba42b..3245cf38 100644 --- a/tutorials/README.md +++ b/tutorials/README.md @@ -34,14 +34,21 @@ current working directory. | 3 | [tutorial_03_lung_reconstruct_highres_4d_ct.py](tutorial_03_lung_reconstruct_highres_4d_ct.py) | `WorkflowReconstructHighres4DCT` | DirLab-4DCT (manual) | | 4 | [tutorial_04_heart_ct_to_vtk.py](tutorial_04_heart_ct_to_vtk.py) | `WorkflowConvertImageToVTK` | Slicer-Heart-CT (prepare first) | | 4 | [tutorial_04_lung_ct_to_vtk.py](tutorial_04_lung_ct_to_vtk.py) | `WorkflowConvertImageToVTK` | Lung gated 4D CT (prepare first) | +| 4 | [duke heart labelmap variant](tutorial_04_duke_heart_labelmap_to_vtk.py) | `ContourTools.extract_label_surfaces`, `ContourTools.extract_tetrahedra` | Duke-Heart-4DLabelmaps (not yet available) | | 5 | [tutorial_05_heart_vtk_to_usd.py](tutorial_05_heart_vtk_to_usd.py) | `WorkflowConvertVTKToUSD` | Output of tutorial 4 | +| 5 | [duke heart variant](tutorial_05_duke_heart_vtk_to_usd.py) | `ConvertVTKToUSD`, `USDAnatomyTools` | Output of tutorial 4 (duke heart labelmap) | | 6 | [tutorial_06_heart_create_statistical_model.py](tutorial_06_heart_create_statistical_model.py) | `WorkflowCreateStatisticalModel` | KCL-Heart-Model | | 6 | [tutorial_06_lung_create_statistical_model.py](tutorial_06_lung_create_statistical_model.py) | `WorkflowCreateStatisticalModel` | Lung surfaces from Tutorial 4 (lung) | +| 6 | [duke heart variant](tutorial_06_duke_heart_create_statistical_model.py) | `WorkflowCreateMeanSurface`, `WorkflowCreateStatisticalModel` | Reference-frame heart surfaces from Tutorial 4 (duke heart labelmap) | | 7 | [tutorial_07_heart_fit_statistical_model_to_patient.py](tutorial_07_heart_fit_statistical_model_to_patient.py) | `WorkflowFitStatisticalModelToPatient` | KCL-Heart-Model plus Tutorial 6 output | | 7 | [tutorial_07_lung_fit_statistical_model_to_patient.py](tutorial_07_lung_fit_statistical_model_to_patient.py) | `WorkflowFitStatisticalModelToPatient` | Chest-CT plus Tutorial 6 (lung) output | +| 7 | [duke heart variant](tutorial_07_duke_heart_fit_statistical_model_to_patient.py) | `WorkflowFitStatisticalModelToPatient` | Duke-Heart-4DLabelmaps plus Tutorial 6 (duke heart) output | | 8 | [tutorial_08_lung_fit_model_to_4d_patients.py](tutorial_08_lung_fit_model_to_4d_patients.py) | `WorkflowFitStatisticalModelToPatient`, `WorkflowReconstructHighres4DCT` | DirLab-4DCT plus Tutorial 6 (lung) and Tutorial 2 output | +| 8 | [duke heart variant](tutorial_08_duke_heart_fit_model_to_4d_patients.py) | `WorkflowFitStatisticalModelToPatient`, `RegisterModelsDistanceMaps` | Duke-Heart-4DLabelmaps plus Tutorial 6 (duke heart) and Tutorial 2 output | | 9 | [tutorial_09_lung_train_physicsnemo_mgn.py](tutorial_09_lung_train_physicsnemo_mgn.py) | `WorkflowTrainPhysicsNeMo`, `WorkflowInferPhysicsNeMo`, `WorkflowInferMovement` (requires `[physicsnemo]` extra + `torch-geometric`) | Tutorial 8 (lung) output | +| 9 | [duke heart variant](tutorial_09_duke_heart_train_physicsnemo_mgn.py) | `WorkflowTrainPhysicsNeMo`, `WorkflowInferPhysicsNeMo`, `WorkflowInferMovement` (requires `[physicsnemo]` extra + `torch-geometric`) | Tutorial 8 (duke heart) output | | 10 | [tutorial_10_lung_infer_physicsnemo_mgn.py](tutorial_10_lung_infer_physicsnemo_mgn.py) | `WorkflowInferPhysicsNeMo`, `WorkflowInferMovement`, `WorkflowConvertVTKToUSD` (requires `[physicsnemo]` extra + `torch-geometric`) | Tutorial 8 and 9 (lung) output | +| 10 | [duke heart variant](tutorial_10_duke_heart_infer_physicsnemo.py) | `WorkflowInferPhysicsNeMo`, `WorkflowInferMovement`, `WorkflowConvertVTKToUSD` (requires `[physicsnemo]` extra + `torch-geometric`) | Tutorial 8 and 9 (duke heart) output | The [tutorials page](https://project-monai.github.io/physiotwin4d/tutorials.html) covers the same set with previews of what each one produces and per-tutorial diff --git a/tutorials/parameters_duke_heart_labelmaps.py b/tutorials/parameters_duke_heart_labelmaps.py new file mode 100644 index 00000000..8319bb87 --- /dev/null +++ b/tutorials/parameters_duke_heart_labelmaps.py @@ -0,0 +1,152 @@ +"""Shared parameters for the Duke-Heart-4DLabelmaps tutorials. + +Mirrors :mod:`parameters_heart_ct_kcl` for the ``duke_heart`` tutorials, which +read segmented labelmaps rather than CT and therefore need their own label ids +and their own held-out case. The registration values match the KCL heart ones +on purpose: Tutorial 2 finetunes uniGradICON on distance maps whose appearance +is fixed by ``mask_dilation_mm`` and ``distancemap_squared_max``, and Tutorial 7 +infers with those weights, so the two must agree. + +The directories and the shape-model files the tutorials read and write live here +too, so that Tutorial 6 writes the model where Tutorial 7 looks for it. Every +path is derived from this file's location, so they hold wherever the clone is. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path + +from physiotwin4d import SegmentAnatomyBase, SegmentHeartSimplewareTrimmedBranches + +_REPO_ROOT = Path(__file__).resolve().parent.parent +_TUTORIAL_OUTPUT = _REPO_ROOT / "tutorials" / "output" +_TUTORIAL_WEIGHTS = _REPO_ROOT / "tutorials" / "network_weights" + + +@dataclass(frozen=True) +class ParametersDukeHeartLabelmaps: + """Settings shared by the Duke heart tutorials. + + Attributes: + mask_dilation_mm: Dilation of the binary registration masks, in + millimeters. + distancemap_squared_max: Saturation radius of every heart distance map, + in squared millimeters. Fixes their intensity distribution, so the + finetuning tutorial and every tutorial that registers these distance + maps must use this one value. + surface_spacing_mm: Isotropic pitch of the grid every heart surface is + contoured on. Finer than the slice pitch of this data, which is + what lets a boundary land between two slices instead of terracing at + one of them. + surface_smoothing_iterations: Taubin smoothing iterations applied to + those surfaces. + mesh_element_size_mm: Edge length of the isotropic voxels a labelmap is + resampled to before it is meshed into tetrahedra, which is the + resulting element size. Below the thinnest wall of the heart, so + that the myocardium survives the coarsening. + model_points: Points kept per surface when building the shape model. + The contours carry thirty times this, which a model built from a few + dozen patients cannot support and every PCA mode would have to + carry as one vector of ``3 * model_points``. + number_of_pca_components: PCA components retained when building the + heart statistical model, and used when fitting it to a patient. + number_of_pca_components_test: Same, under ``TestTools.running_as_test``. + number_of_iterations_greedy: Greedy coarse-to-fine iteration schedule. + number_of_iterations_greedy_test: Same, under + ``TestTools.running_as_test``. + segmenter_class: Segmenter that produced these labelmaps, so the + tutorials name their labels the way it does. + anatomy_group: Anatomy group name that segmenter registers for the heart. + interior_object_ids: Labels left out of the whole-heart structure, and + therefore never measured to by a distance map. The four chambers + (1-4) cover the cavities rather than their walls, so measuring to + them measures the inside of the heart; the great vessels and + coronaries (7-10) vary too much in extent between patients to be + part of a shape model. What remains is the myocardium (5) and the + heart wall (6). + input_dir: Surfaces Tutorial 6 builds the model from, which are what + Tutorial 4 wrote; ``input_dir_test`` is its counterpart under + ``TestTools.running_as_test``. + hold_out_dir: Labelmaps the held-out case is read from by Tutorials 2 + and 7, and ``hold_out_dir_test`` its counterpart under + ``TestTools.running_as_test``. + pca_json_file: Shape model Tutorial 6 writes and Tutorials 7, 8 and 9 + read. + pca_mean_file: Mean surface of that model, written and read the same way. + mgn_weights_dir: MeshGraphNet Tutorial 9 trains and Tutorial 10 infers + with, alongside the normalization statistics and PCA assets that + make that directory self-contained. It sits beside the finetuned + ICON checkpoints rather than under ``output/``, because it is a + trained network rather than a per-run result. + hold_out_case: Case held out of every fit: Tutorial 6 builds the shape + model without it and Tutorial 7 fits that model to it, so the fit + measures generalization rather than reconstruction. Tutorial 2 + scores its registrations on the same case. + """ + + mask_dilation_mm: float = 10.0 + distancemap_squared_max: float = (1.25 * 10.0) ** 2 + + surface_spacing_mm: float = 0.5 + surface_smoothing_iterations: int = 20 + mesh_element_size_mm: float = 1.5 + + model_points: int = 20000 + number_of_pca_components: int = 10 + number_of_pca_components_test: int = 5 + + number_of_iterations_greedy: list[int] = field( + default_factory=lambda: [30, 15, 7, 3] + ) + number_of_iterations_greedy_test: list[int] = field(default_factory=lambda: [1, 0]) + + segmenter_class: type[SegmentAnatomyBase] = SegmentHeartSimplewareTrimmedBranches + anatomy_group: str = "heart" + interior_object_ids: list[int] = field( + default_factory=lambda: [1, 2, 3, 4, 7, 8, 9, 10] + ) + + hold_out_case: str = "pm0027" + + input_dir: Path = _TUTORIAL_OUTPUT / "tutorial_04_duke_heart_labelmap" + input_dir_test: Path = _TUTORIAL_OUTPUT / "tutorial_04_duke_heart_labelmap" + hold_out_dir: Path = _REPO_ROOT / "data" / "Duke-Heart-4DLabelmaps" + hold_out_dir_test: Path = _REPO_ROOT / "data" / "test" / "Duke-Heart-4DLabelmaps" + pca_json_file: Path = _TUTORIAL_OUTPUT / "tutorial_06_duke_heart" / "pca_model.json" + pca_mean_file: Path = ( + _TUTORIAL_OUTPUT / "tutorial_06_duke_heart" / "pca_mean_surface.vtp" + ) + mgn_weights_dir: Path = _TUTORIAL_WEIGHTS / "physicsnemo_mgn_duke_heart_motion" + + def input_directory(self, test_mode: bool) -> Path: + """Return the model population directory for this run mode. + + Tutorial 4 writes to one directory whichever mode it ran in, so both + modes read the same one. + """ + return self.input_dir_test if test_mode else self.input_dir + + def hold_out_directory(self, test_mode: bool) -> Path: + """Return the held-out case's directory for this run mode.""" + return self.hold_out_dir_test if test_mode else self.hold_out_dir + + def pca_components(self, test_mode: bool) -> int: + """Return the PCA component count for this run mode.""" + return ( + self.number_of_pca_components_test + if test_mode + else self.number_of_pca_components + ) + + def greedy_iterations(self, test_mode: bool) -> list[int]: + """Return the Greedy iteration schedule for this run mode.""" + return list( + self.number_of_iterations_greedy_test + if test_mode + else self.number_of_iterations_greedy + ) + + +#: The single instance every Duke heart tutorial imports. +DUKE_HEART = ParametersDukeHeartLabelmaps() diff --git a/tutorials/parameters_heart_ct_kcl.py b/tutorials/parameters_heart_ct_kcl.py index 4d4bfa14..1b76c3e6 100644 --- a/tutorials/parameters_heart_ct_kcl.py +++ b/tutorials/parameters_heart_ct_kcl.py @@ -7,16 +7,21 @@ reusing the lung one's weights -- the two organs' distance maps do not share an intensity distribution. -Paths stay out of this module: each tutorial owns its own inputs and outputs. +The directories and the shape-model files the tutorials read and write live here +too, so that Tutorial 6 writes the model where Tutorial 7 looks for it. Every +path is derived from this file's location, so they hold wherever the clone is. """ from __future__ import annotations from dataclasses import dataclass, field -from typing import Optional +from pathlib import Path from physiotwin4d import SegmentAnatomyBase, SegmentChestTotalSegmentator +_REPO_ROOT = Path(__file__).resolve().parent.parent +_OUTPUT_ROOT = _REPO_ROOT / "tutorials" / "output" / "tutorial_06_heart" + @dataclass(frozen=True) class ParametersHeartCTKCL: @@ -30,6 +35,12 @@ class ParametersHeartCTKCL: in squared millimeters. Fixes their intensity distribution, so the finetuning tutorial and every tutorial that registers heart distance maps must use this one value. + surface_reduction_rate: Fraction of triangles removed from every + extracted heart surface. ``0.0`` keeps them at full resolution. + mesh_element_size_mm: Edge length of the isotropic voxels a labelmap + is resampled to before it is meshed into tetrahedra, which is the + resulting element size. Below the thinnest wall of the heart, so + that the myocardium survives the coarsening. number_of_pca_components: PCA components retained when building the heart statistical model, and used when fitting it to a patient. number_of_pca_components_test: Same, under ``TestTools.running_as_test``. @@ -42,15 +53,31 @@ class ParametersHeartCTKCL: interior_object_ids_totalsegmentator: Chamber labels in a TotalSegmentator labelmap. The chambers are interior to the myocardium, so a distance map must not measure to them. - interior_object_ids_simpleware: The same chambers in a Simpleware - ASCardio labelmap, which the Duke-Heart-4DLabelmaps data was - segmented with. Which list applies is a property of the data a - tutorial reads, not of any one segmenter class, so both live here. + input_dir: Population Tutorial 6 builds the model from, and + ``input_dir_test`` its counterpart under + ``TestTools.running_as_test``. + hold_out_dir: Dataset the held-out case is read from by Tutorial 7, and + ``hold_out_dir_test`` its counterpart under + ``TestTools.running_as_test``. A different dataset from + ``input_dir``: the model is built from KCL meshes and fitted to a + DIR-Lab patient. + pca_json_file: Shape model Tutorial 6 writes and Tutorial 7 reads. + pca_mean_file: Mean surface of that model, written and read the same way. + hold_out_case: DIR-Lab case fitted by Tutorial 7 and therefore kept out + of the population Tutorial 6 builds the model from, so that the fit + measures generalization rather than reconstruction. The KCL model + meshes carry no DIR-Lab case, so today the exclusion never fires; + Tutorial 6 applies it anyway, so adding one cannot slip it in. + The Duke heart tutorials name their own in + ``parameters_duke_heart_labelmaps.py``. """ mask_dilation_mm: float = 10.0 distancemap_squared_max: float = (1.25 * 10.0) ** 2 + surface_reduction_rate: float = 0.5 + mesh_element_size_mm: float = 1.5 + number_of_pca_components: int = 10 number_of_pca_components_test: int = 5 @@ -61,12 +88,26 @@ class ParametersHeartCTKCL: segmenter_class: type[SegmentAnatomyBase] = SegmentChestTotalSegmentator anatomy_group: str = "heart" - interior_object_ids_totalsegmentator: Optional[list[int]] = field( + interior_object_ids_totalsegmentator: list[int] = field( default_factory=lambda: [141, 142, 143, 144] ) - interior_object_ids_simpleware: list[int] = field( - default_factory=lambda: [1, 2, 3, 4] - ) + + hold_out_case: str = "Case1Pack" + + input_dir: Path = _REPO_ROOT / "data" / "KCL-Heart-Model" + input_dir_test: Path = _REPO_ROOT / "data" / "test" / "KCL-Heart-Model" + hold_out_dir: Path = _REPO_ROOT / "data" / "DirLab-4DCT" + hold_out_dir_test: Path = _REPO_ROOT / "data" / "test" / "DirLab-4DCT" + pca_json_file: Path = _OUTPUT_ROOT / "pca_model.json" + pca_mean_file: Path = _OUTPUT_ROOT / "pca_mean_surface.vtp" + + def input_directory(self, test_mode: bool) -> Path: + """Return the model population directory for this run mode.""" + return self.input_dir_test if test_mode else self.input_dir + + def hold_out_directory(self, test_mode: bool) -> Path: + """Return the held-out case's directory for this run mode.""" + return self.hold_out_dir_test if test_mode else self.hold_out_dir def pca_components(self, test_mode: bool) -> int: """Return the PCA component count for this run mode.""" diff --git a/tutorials/parameters_lung_ct_dirlab.py b/tutorials/parameters_lung_ct_dirlab.py index 08cbb684..438beec4 100644 --- a/tutorials/parameters_lung_ct_dirlab.py +++ b/tutorials/parameters_lung_ct_dirlab.py @@ -6,15 +6,22 @@ dilation that drifts between two of these scripts silently trains on one image distribution and infers on another. -Paths stay out of this module: each tutorial owns its own inputs and outputs. +The directories and the shape-model files the tutorials read and write live here +too, so that Tutorial 6 writes the model where Tutorials 7 and 8 look for it. +Every path is derived from this file's location, so they hold wherever the clone +is. """ from __future__ import annotations from dataclasses import dataclass, field +from pathlib import Path from physiotwin4d import SegmentAnatomyBase, SegmentNVSegmentCTMRI +_REPO_ROOT = Path(__file__).resolve().parent.parent +_OUTPUT_ROOT = _REPO_ROOT / "tutorials" / "output" / "tutorial_06_lung" + @dataclass(frozen=True) class ParametersLungCTDirLab: @@ -28,6 +35,11 @@ class ParametersLungCTDirLab: in squared millimeters. Fixes their intensity distribution, so the finetuning tutorial and every tutorial that registers lung distance maps must use this one value. + surface_reduction_rate: Fraction of triangles removed from every + extracted lung surface. ``0.0`` keeps them at full resolution. + mesh_reduction_rate: Fraction of the voxel resolution removed along + each axis before a labelmap is meshed into tetrahedra, so the + tetrahedron count falls by roughly ``(1 - rate) ** 3``. number_of_pca_components: PCA components retained when building the lung statistical model, and used when fitting it to a patient. number_of_pca_components_test: Same, under ``TestTools.running_as_test``. @@ -37,6 +49,34 @@ class ParametersLungCTDirLab: segmenter_class: Segmenter every lung tutorial instantiates, so the surfaces they compare share a definition of "lung". anatomy_group: Anatomy group name that segmenter registers for lungs. + input_dir: Population Tutorial 6 builds the model from, and + ``input_dir_test`` its counterpart under + ``TestTools.running_as_test``. + hold_out_dir: Dataset the held-out study is read from by Tutorial 7, and + ``hold_out_dir_test`` its counterpart under + ``TestTools.running_as_test``. + pca_json_file: Shape model Tutorial 6 writes and Tutorials 7 and 8 read. + pca_mean_file: Mean surface of that model, written and read the same way. + hold_out_case: Image fitted by Tutorial 7 and therefore kept out of the + population Tutorial 6 builds the model from, so that the fit + measures generalization rather than reconstruction. It is a Chest-CT + study while the model is built from DIR-Lab phases, so today the + exclusion never fires; Tutorial 6 applies it anyway, so adding the + study to that population cannot slip it in. Tutorial 2 holds out a + DIR-Lab case of its own, which measures registration rather than + shape. + mgn_weights_dir: Directory the lung-motion MeshGraphNet is trained into + by Tutorial 9 and loaded from by Tutorial 10, beside the ICON + weights the registration tutorials finetune. Tutorial 9 writes to a + numbered sibling of it when resuming from a checkpoint, in which + case Tutorial 10 has to be pointed at that sibling. + mgn_hold_out_case: DIR-Lab case kept out of the Tutorial 9 training and + predicted by Tutorial 10, so that the prediction measures + generalization. Distinct from ``hold_out_case``, which is the + Chest-CT study the shape model is fitted to: this one is a 4D case + Tutorial 8 has fitted every phase of. It is also the case Tutorial + 2 holds out of its ICON finetuning, so the surfaces it is scored on + came from a registration network that never saw it either. There is no interior-structure list here, the counterpart of the heart's chamber ids: the lung labels are the lobes, and every one of them is on the @@ -46,6 +86,9 @@ class ParametersLungCTDirLab: mask_dilation_mm: float = 40.0 distancemap_squared_max: float = (1.25 * 40.0) ** 2 + surface_reduction_rate: float = 0.0 + mesh_reduction_rate: float = 0.0 + number_of_pca_components: int = 6 number_of_pca_components_test: int = 5 @@ -57,6 +100,27 @@ class ParametersLungCTDirLab: segmenter_class: type[SegmentAnatomyBase] = SegmentNVSegmentCTMRI anatomy_group: str = "lung" + hold_out_case: str = "Chest-CT.mha" + mgn_hold_out_case: str = "Case1Pack" + + input_dir: Path = _REPO_ROOT / "data" / "DirLab-4DCT" + input_dir_test: Path = _REPO_ROOT / "data" / "test" / "DirLab-4DCT" + hold_out_dir: Path = _REPO_ROOT / "data" / "Chest-CT" + hold_out_dir_test: Path = _REPO_ROOT / "data" / "test" / "Chest-CT" + pca_json_file: Path = _OUTPUT_ROOT / "pca_model.json" + pca_mean_file: Path = _OUTPUT_ROOT / "pca_mean_surface.vtp" + mgn_weights_dir: Path = ( + _REPO_ROOT / "tutorials" / "network_weights" / "physicsnemo_mgn_lung_motion" + ) + + def input_directory(self, test_mode: bool) -> Path: + """Return the model population directory for this run mode.""" + return self.input_dir_test if test_mode else self.input_dir + + def hold_out_directory(self, test_mode: bool) -> Path: + """Return the held-out study's directory for this run mode.""" + return self.hold_out_dir_test if test_mode else self.hold_out_dir + def pca_components(self, test_mode: bool) -> int: """Return the PCA component count for this run mode.""" return ( diff --git a/tutorials/tutorial_01_heart_gated_ct_to_usd.py b/tutorials/tutorial_01_heart_gated_ct_to_usd.py index dcaa4b72..ab70e580 100644 --- a/tutorials/tutorial_01_heart_gated_ct_to_usd.py +++ b/tutorials/tutorial_01_heart_gated_ct_to_usd.py @@ -71,6 +71,7 @@ from pathlib import Path import itk +from parameters_heart_ct_kcl import HEART_CT_KCL from physiotwin4d import ( RegisterImagesGreedy, @@ -140,6 +141,7 @@ usd_project_name="cardiac_model", registration_method=registration_method, segmentation_method=segmentation_method, + surface_reduction_rate=HEART_CT_KCL.surface_reduction_rate, log_level=log_level, save_assets=True, ) diff --git a/tutorials/tutorial_01_lung_gated_ct_to_usd.py b/tutorials/tutorial_01_lung_gated_ct_to_usd.py index d34e7f3c..918bda65 100644 --- a/tutorials/tutorial_01_lung_gated_ct_to_usd.py +++ b/tutorials/tutorial_01_lung_gated_ct_to_usd.py @@ -69,6 +69,7 @@ from pathlib import Path import itk +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( RegisterImagesGreedy, @@ -140,6 +141,7 @@ usd_project_name="lung_model", registration_method=registration_method, segmentation_method=segmentation_method, + surface_reduction_rate=LUNG_CT_DIRLAB.surface_reduction_rate, log_level=log_level, frames_per_second=1, save_assets=True, diff --git a/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py b/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py index 34658b4a..16ae3668 100644 --- a/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py +++ b/tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py @@ -12,11 +12,11 @@ The heart needs its own run rather than reusing the lung weights. A distance map's appearance is set by the radius it saturates at, and the heart is registered with a far tighter mask than the lungs -(``ParametersHeartCTKCL.mask_dilation_mm`` versus +(``ParametersDukeHeartLabelmaps.mask_dilation_mm`` versus ``ParametersLungCTDirLab.mask_dilation_mm``), so the two organs' distance maps do not share an intensity distribution. Every value that fixes that -distribution comes from ``parameters_heart_ct_kcl.py``, so this run trains on -exactly what Tutorial 7 (heart) later infers on. +distribution comes from ``parameters_duke_heart_labelmaps.py``, so this run +trains on exactly what Tutorial 7 (Duke Heart) later infers on. Duke-Heart-4DLabelmaps ships segmented labelmaps, not CT, which is all this tutorial needs. Each labelmap serves twice: with the chamber-interior labels @@ -25,10 +25,12 @@ ``SegmentHeartSimplewareTrimmedBranches``, whose ventricle and atrium labels cover the cavities rather than the walls, so measuring distance to them would measure distance to the inside of the heart; the ids to drop come from -``parameters_heart_ct_kcl.py`` alongside every other value this data needs. +``parameters_duke_heart_labelmaps.py`` alongside every other value this data +needs. -The first patient is held out and two of its gated phases are registered three -ways -- Greedy on the distance maps, stock ICON, and the finetuned ICON -- and +``ParametersDukeHeartLabelmaps.hold_out_case`` is held out -- the same case +Tutorials 6 and 7 keep out of the shape model -- and two of its gated phases are +registered three ways -- Greedy on the distance maps, stock ICON, and the finetuned ICON -- and scored by target registration error over the anatomical landmarks Slicer markups files supply for every frame, plus per-class Dice. @@ -52,7 +54,7 @@ import itk import numpy as np -from parameters_heart_ct_kcl import HEART_CT_KCL +from parameters_duke_heart_labelmaps import DUKE_HEART from physiotwin4d import ( ContourTools, @@ -90,22 +92,21 @@ # Distance-map normalization, shared with every heart tutorial that later # registers these maps, so this run trains on the same image distribution # they infer on. - distance_squared_max = HEART_CT_KCL.distancemap_squared_max + distance_squared_max = DUKE_HEART.distancemap_squared_max # Set True to finetune from scratch. That deletes experiment_dir below, # including any checkpoint a previous run left there. run_finetuning = True test_mode = TestTools.running_as_test() + data_dir = DUKE_HEART.hold_out_directory(test_mode) if test_mode: - data_dir = repo_root / "data" / "test" / "Duke-Heart-4DLabelmaps" number_of_iterations_icon: Optional[int] = 1 epochs = 1 else: - data_dir = repo_root / "data" / "Duke-Heart-4DLabelmaps" number_of_iterations_icon = 10 epochs = 100 - number_of_iterations_greedy = HEART_CT_KCL.greedy_iterations(test_mode) + number_of_iterations_greedy = DUKE_HEART.greedy_iterations(test_mode) log_level = logging.INFO reporter = PhysioTwin4DBase(class_name=class_name, log_level=log_level) @@ -115,9 +116,9 @@ contour_tools = ContourTools(log_level=log_level) transform_tools = TransformTools() - # This data was segmented with Simpleware ASCardio, so the chamber labels - # to drop are that labelmap's, not TotalSegmentator's. - interior_object_ids = HEART_CT_KCL.interior_object_ids_simpleware + # Labels left out of the surface a distance map is measured to; see the + # parameters module. + interior_object_ids = DUKE_HEART.interior_object_ids # Cohort discovery case_dirs = sorted(path for path in data_dir.glob("pm*") if path.is_dir()) @@ -146,7 +147,11 @@ def distance_map_for(labelmap_file: Path) -> Path: [-1000, 1000] window uniGradICON's CT preprocessing expects. """ stem = labelmap_file.name[: -len("_labelmap.nii.gz")] - distance_map_file = derived_dir / f"{stem}_distance_map.mha" + # Frame stems repeat across cases, so the case directory name is part of + # the cache key; without it one case's map would be read for another's. + distance_map_file = ( + derived_dir / f"{labelmap_file.parent.name}_{stem}_distance_map.mha" + ) if distance_map_file.exists(): return distance_map_file @@ -177,9 +182,13 @@ def distance_map_for(labelmap_file: Path) -> Path: itk.imwrite(distance_map, str(distance_map_file), compression=True) return distance_map_file - # Held-out patient: the first case, excluded from finetuning entirely. - held_out_dir = case_dirs[0] - training_dirs = case_dirs[1:] + # Held-out patient, excluded from finetuning entirely: the case Tutorials 6 + # and 7 also keep out, so one patient is unseen by everything downstream. + held_out_dir = next( + (path for path in case_dirs if path.name == DUKE_HEART.hold_out_case), + case_dirs[0], + ) + training_dirs = [path for path in case_dirs if path != held_out_dir] held_out_frames = frames_for_case(held_out_dir) if len(held_out_frames) < 2: @@ -209,6 +218,11 @@ def distance_map_for(labelmap_file: Path) -> Path: [str(distance_map_for(frame)) for frame in frames] ) subject_labelmap_files.append([str(frame) for frame in frames]) + if not subject_ids: + raise FileNotFoundError( + f"No training case under {data_dir} has the 2+ frames paired " + "finetuning needs; every case but the held-out one was skipped." + ) reporter.log_info( "Finetuning cohort: %d cases, %d frames (held out %s)", len(subject_ids), @@ -277,11 +291,20 @@ def read_landmarks(labelmap_file: Path) -> dict[str, np.ndarray]: """Read a frame's Slicer markups file as ``{label: LPS point}``. The markups files declare ``coordinateSystem: LPS``, the frame this - project works in, so the control points are used as written. + project works in, so the control points are used as written. A file + written in Slicer's RAS default would flip X and Y, so the declaration + is checked rather than assumed. """ landmark_file = companion(labelmap_file, "_landmark.mrk.json") with landmark_file.open(encoding="utf-8") as f: markups = json.load(f)["markups"] + for markup in markups: + coordinate_system = markup.get("coordinateSystem") + if coordinate_system != "LPS": + raise ValueError( + f"{landmark_file.name} declares coordinateSystem " + f"{coordinate_system!r}; this tutorial reads LPS markups." + ) return { point["label"]: np.asarray(point["position"], dtype=np.float64) for markup in markups diff --git a/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py b/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py new file mode 100644 index 00000000..e7340afb --- /dev/null +++ b/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py @@ -0,0 +1,356 @@ +""" +Tutorial 4 (Duke Heart): 4D Labelmaps to VTK Surfaces and Tetrahedral Meshes + +Purpose +------- +Turn the Duke-Heart-4DLabelmaps labelmaps into VTK geometry. ``outputs`` +below chooses how much of it to build, because the full pass costs about a +hundred times the short one: + +``"shape_model"`` + Each case's reference frame only (``*_ref_labelmap.nii.gz``), and only the + whole heart, as ``_heart_minus_interior_chambers.vtp`` and its + ``.vtu`` tetrahedral mesh. Tutorial 6 (Duke Heart) reads the surfaces to + build its shape model. + +``"full"`` + Every gated frame and every label, each with its own surface and mesh. + Tutorial 5 (Duke Heart) needs this, since it animates the frames. + +Files written per frame, the per-label ones only under ``"full"``: + +- ``_surfaces.vtp`` -- every label's watertight, outward-oriented + surface in one file, via ``ContourTools.extract_label_surfaces`` on the whole + labelmap. Extracting the labels together is what keeps neighbors touching: + a wall between two of them is contoured from the same field on the same + isotropic grid, so both surfaces carry the same vertices there. Per-cell + ``SegmentationLabelIds`` says which label each triangle came from. +- ``_.vtu`` -- one tetrahedral mesh per structure, six + ``VTK_TETRA`` per isotropic voxel, via ``ContourTools.extract_tetrahedra``, + then relaxed onto that structure's surface by + ``ContourTools.trim_tetrahedra_to_surface``. + +The mesh starts as a voxel staircase and ends up bounded by the smooth +surface: the relaxation projects its boundary onto that surface while smoothing +the interior to make room, so no trace of the voxel blocks is left. What +remains between the two geometries is faceting at the element size, logged per +structure in millimeters. + +Both carry the structure's ``USDAnatomyTools`` color, as ``AnatomyColor`` and +as a per-cell ``Color``, so they render the same way as the surfaces +``WorkflowConvertImageToVTK`` writes. + +The whole-heart structure drops the labels ``interior_object_ids`` names in +``parameters_duke_heart_labelmaps.py`` -- the chamber cavities, and the vessels +whose extent varies too much between patients -- which is the same definition of +"heart" that Tutorials 2, 6 and 7 use. + +Label names come from ``SegmentHeartSimplewareTrimmedBranches``'s taxonomy, +the segmenter that produced these labelmaps. Labels with no taxonomy entry are +named ``label_``. + +Data Required +------------- +``data/Duke-Heart-4DLabelmaps/pm????/*_labelmap.nii.gz`` (multi-label) +""" + +# Imports +from __future__ import annotations + +import logging +from pathlib import Path +from typing import Optional + +import itk +import numpy as np +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + ContourTools, + PhysioTwin4DBase, + SegmentHeartSimplewareTrimmedBranches, + TestTools, +) + +# Only run if this script is not imported as a module +if __name__ == "__main__": + # Data directory specification + repo_root = Path(__file__).resolve().parent.parent + tutorials_dir = Path(__file__).resolve().parent + + class_name = "tutorial_04_duke_heart_labelmap_to_vtk" + + output_dir = tutorials_dir / "output" / "tutorial_04_duke_heart_labelmap" + + # How much to build; see this module's docstring. "shape_model" is what + # Tutorial 6 (Duke Heart) consumes and runs in minutes over the cohort; + # "full" also feeds Tutorial 5 (Duke Heart) and takes hours, because it + # contours and meshes every label of all 348 gated frames rather than the + # whole heart of the 29 reference ones. + outputs = "shape_model" + if outputs not in ("shape_model", "full"): + raise ValueError(f"outputs must be 'shape_model' or 'full'; got {outputs!r}") + reference_frames_only = outputs == "shape_model" + write_per_label_surfaces = outputs == "full" + + # Taubin smoothing iterations applied to the extracted surfaces (0 disables). + smoothing_iterations = DUKE_HEART.surface_smoothing_iterations + + # Surfaces are contoured on an isotropic grid of this pitch, which also sets + # the triangle count in place of decimation: decimating a label would move + # its vertices off the neighbor it shares them with. + surface_spacing_mm = DUKE_HEART.surface_spacing_mm + + # Element size of the tetrahedral meshes. Isotropic, so the elements do + # not inherit the slice pitch, and small enough that relaxing them onto the + # surface leaves no trace of the voxel staircase. + mesh_element_size_mm = DUKE_HEART.mesh_element_size_mm + + # Labels left out of the whole-heart structure; see the parameters module. + interior_object_ids = DUKE_HEART.interior_object_ids + whole_heart_name = "heart_minus_interior_chambers" + + # Set False to keep outputs a previous run already wrote. + overwrite = True + + test_mode = TestTools.running_as_test() + if test_mode: + data_dir = repo_root / "data" / "test" / "Duke-Heart-4DLabelmaps" + else: + data_dir = repo_root / "data" / "Duke-Heart-4DLabelmaps" + + log_level = logging.INFO + reporter = PhysioTwin4DBase(class_name=class_name, log_level=log_level) + + contour_tools = ContourTools(log_level=log_level) + taxonomy = SegmentHeartSimplewareTrimmedBranches(log_level=logging.WARNING).taxonomy + label_names = taxonomy.all_labels() + + def mask_from(labelmap_image: itk.Image, keep: np.ndarray) -> itk.Image: + """Return *keep* as a binary mask carrying *labelmap_image*'s geometry.""" + mask = itk.GetImageFromArray(keep.astype(np.uint8)) + mask.CopyInformation(labelmap_image) + return mask + + def annotate( + mesh: pv.DataSet, label_ids: list[int], name: str, source: str + ) -> None: + """Record the originating labels and file on *mesh* in-place. + + ``SegmentationLabelIds`` holds one id for a per-label structure and + every retained id for the whole-heart one; a single id is what + ``ContourTools.save_combined_surfaces`` needs to tag the merged file's + cells with the structure they came from. The anatomy color is attached + separately, by ``ContourTools.apply_anatomy_color``. + """ + mesh.field_data["SegmentationLabelIds"] = np.asarray(label_ids, dtype=np.int32) + mesh.field_data["LabelName"] = np.array([name]) + mesh.field_data["SourceLabelmap"] = np.array([source]) + + def surface_to_mesh_displacement( + surface: pv.PolyData, mesh: pv.UnstructuredGrid + ) -> float: + """Return the mean distance, in mm, from *surface*'s points to *mesh*.""" + boundary = mesh.extract_surface(algorithm="dataset_surface") + # On a copy, so the distances are not written into the saved surface. + distances = surface.copy().compute_implicit_distance(boundary) + return float(np.abs(distances["implicit_distance"]).mean()) + + def write_tetrahedra( + mask: itk.Image, + surface: pv.PolyData, + case_output_dir: Path, + stem: str, + name: str, + label_ids: list[int], + source: str, + ) -> Optional[float]: + """Write one structure's VTU mesh, relaxed onto its *surface*. + + Returns: + The mean displacement, in millimeters, between the two geometries, + or ``None`` for a structure too thin to hold an element. + """ + # Resampling to the element size eats a structure thinner than it -- + # the coronary arteries here are a voxel or two across -- so the size + # is halved until the mesh accounts for most of the surface's volume, + # or until it is as fine as the labelmap itself. + finest_spacing = float(np.min(np.asarray(mask.GetSpacing()))) + element_size = mesh_element_size_mm + while True: + # USDAnatomyTools has no override for names like "left_ventricle", + # so the structure's anatomy group is offered as the fallback color. + tet_mesh = contour_tools.extract_tetrahedra( + mask, + element_size_mm=element_size, + anatomy_names=[name, taxonomy.group_for_label(name)], + ) + volume = float(np.sum(tet_mesh.compute_cell_sizes(volume=True)["Volume"])) + if volume > 0.5 * surface.volume or element_size <= finest_spacing: + break + element_size *= 0.5 + reporter.log_debug( + " %s: %.3g mm elements hold %.0f%% of it; halving", + name, + element_size * 2.0, + 100.0 * volume / float(surface.volume), + ) + if tet_mesh.n_cells == 0: + reporter.log_warning(" %-24s thinner than one element; no mesh", name) + return None + tet_mesh = contour_tools.trim_tetrahedra_to_surface(tet_mesh, surface) + annotate(tet_mesh, label_ids, name, source) + tet_mesh.save(case_output_dir / f"{stem}_{name}.vtu") + + displacement = surface_to_mesh_displacement(surface, tet_mesh) + reporter.log_info( + " %-24s %7d triangles %8d tetrahedra %6.3f mm", + name, + surface.n_cells, + tet_mesh.n_cells, + displacement, + ) + return displacement + + # Cohort discovery + case_dirs = sorted( + path for path in data_dir.glob("pm[0-9][0-9][0-9][0-9]") if path.is_dir() + ) + if not case_dirs: + raise FileNotFoundError( + f"No pm???? case directories found under {data_dir}.\n" + "See data/README.md for download instructions." + ) + + # Extraction + labelmap_pattern = ( + "*_ref_labelmap.nii.gz" if reference_frames_only else "*_labelmap.nii.gz" + ) + displacements: list[float] = [] + surface_count = 0 + whole_heart_surface: Optional[pv.PolyData] = None + + case_output_dir = output_dir + case_output_dir.mkdir(parents=True, exist_ok=True) + + for case_dir in case_dirs: + labelmap_files = sorted(case_dir.glob(labelmap_pattern)) + reporter.log_section(f"{case_dir.name}: {len(labelmap_files)} labelmaps") + + for labelmap_file in labelmap_files: + stem = labelmap_file.name[: -len("_labelmap.nii.gz")] + surfaces_file = case_output_dir / f"{stem}_surfaces.vtp" + whole_heart_file = case_output_dir / f"{stem}_{whole_heart_name}.vtp" + # The whole-heart surface is written last, so its presence means + # the frame finished rather than stopped part way. + if not overwrite and whole_heart_file.exists(): + reporter.log_debug("%s: outputs exist, skipping", labelmap_file.name) + continue + + reporter.log_info("%s", labelmap_file.name) + labelmap_image = itk.imread(str(labelmap_file)) + labels = itk.GetArrayViewFromImage(labelmap_image) + present_ids = [int(value) for value in np.unique(labels) if value != 0] + + # The whole heart, hollowed of the chamber cavities. It overlaps + # the per-label surfaces, so it is kept out of their file. + whole_heart_ids = [ + label_id + for label_id in present_ids + if label_id not in interior_object_ids + ] + whole_heart_mask = mask_from( + labelmap_image, np.isin(labels, whole_heart_ids) + ) + whole_heart = contour_tools.extract_label_surfaces( + whole_heart_mask, + isotropic_spacing_mm=surface_spacing_mm, + smoothing_iterations=smoothing_iterations, + )[1] + contour_tools.apply_anatomy_color(whole_heart, [whole_heart_name, "heart"]) + annotate(whole_heart, whole_heart_ids, whole_heart_name, labelmap_file.name) + surface_count += 1 + displacement = write_tetrahedra( + whole_heart_mask, + whole_heart, + case_output_dir, + stem, + whole_heart_name, + whole_heart_ids, + labelmap_file.name, + ) + if displacement is not None: + displacements.append(displacement) + + if write_per_label_surfaces: + # Every label at once, so that neighbors share the wall between + # them rather than each contouring its own copy of it. + surfaces = contour_tools.extract_label_surfaces( + labelmap_image, + isotropic_spacing_mm=surface_spacing_mm, + smoothing_iterations=smoothing_iterations, + ) + named_surfaces: dict[str, pv.PolyData] = {} + for label_id, surface in surfaces.items(): + name = label_names.get(label_id, f"label_{label_id}") + contour_tools.apply_anatomy_color( + surface, [name, taxonomy.group_for_label(name)] + ) + annotate(surface, [label_id], name, labelmap_file.name) + named_surfaces[name] = surface + surface_count += 1 + displacement = write_tetrahedra( + mask_from(labelmap_image, labels == label_id), + surface, + case_output_dir, + stem, + name, + [label_id], + labelmap_file.name, + ) + if displacement is not None: + displacements.append(displacement) + ContourTools.save_combined_surfaces(named_surfaces, str(surfaces_file)) + + # Written last, so an interrupted frame is redone rather than + # skipped by the check above. + whole_heart.save(whole_heart_file) + if whole_heart_surface is None: + whole_heart_surface = whole_heart + + mean_displacement = float(np.mean(displacements)) if displacements else 0.0 + if displacements: + reporter.log_section( + f"Wrote {surface_count} surfaces and {len(displacements)} tetrahedral " + f"meshes, mean surface-to-mesh displacement {mean_displacement:.3f} mm" + ) + else: + reporter.log_section(f"Wrote {surface_count} surfaces") + + # Testing + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + log_level=log_level, + ) + + screenshots: list[Path] = [] + if whole_heart_surface is not None: + screenshots.append( + tt.save_screenshot_mesh( + whole_heart_surface, + f"{whole_heart_name}.png", + camera_position="iso", + color="lightblue", + opacity=0.85, + ) + ) + + tutorial_results = { + "output_dir": output_dir, + "case_dirs": case_dirs, + "n_pairs": len(displacements), + "mean_surface_to_mesh_displacement_mm": mean_displacement, + "screenshots": screenshots, + } diff --git a/tutorials/tutorial_04_heart_ct_to_vtk.py b/tutorials/tutorial_04_heart_ct_to_vtk.py index 00f5cec9..c675c6c9 100644 --- a/tutorials/tutorial_04_heart_ct_to_vtk.py +++ b/tutorials/tutorial_04_heart_ct_to_vtk.py @@ -21,6 +21,7 @@ import itk import pyvista as pv +from parameters_heart_ct_kcl import HEART_CT_KCL from physiotwin4d import ( ContourTools, @@ -102,10 +103,10 @@ # Workflow execution # - # surface_target_reduction decimates each exported VTP surface. + # surface_reduction_rate decimates each exported VTP surface. result = workflow.process( input_image=ct_image, - surface_target_reduction=0.5, + surface_reduction_rate=HEART_CT_KCL.surface_reduction_rate, extract_label_surfaces=save_label_surfaces, ) diff --git a/tutorials/tutorial_04_lung_ct_to_vtk.py b/tutorials/tutorial_04_lung_ct_to_vtk.py index be1738c6..e8bbdf80 100644 --- a/tutorials/tutorial_04_lung_ct_to_vtk.py +++ b/tutorials/tutorial_04_lung_ct_to_vtk.py @@ -22,6 +22,8 @@ import itk import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB + from physiotwin4d import ( ContourTools, SegmentChestTotalSegmentator, @@ -86,10 +88,10 @@ # Workflow execution # - # surface_target_reduction decimates each exported VTP surface. + # surface_reduction_rate decimates each exported VTP surface. result = workflow.process( input_image=ct_image, - surface_target_reduction=0.5, + surface_reduction_rate=LUNG_CT_DIRLAB.surface_reduction_rate, extract_label_surfaces=save_label_surfaces, ) diff --git a/tutorials/tutorial_05_duke_heart_vtk_to_usd.py b/tutorials/tutorial_05_duke_heart_vtk_to_usd.py new file mode 100644 index 00000000..6c1eebc5 --- /dev/null +++ b/tutorials/tutorial_05_duke_heart_vtk_to_usd.py @@ -0,0 +1,155 @@ +""" +Tutorial 5 (Duke Heart): 4D VTK Surfaces to Animated USD + +Purpose +------- +Assemble the per-frame, per-structure surfaces Tutorial 4 (Duke Heart) wrote +into one animated USD per case: every gated frame becomes a time sample, every +label its own prim under ``/World///``, painted +with that structure's OmniSurface material. + +Tutorial 4 saves one VTP per frame holding every structure, each cell tagged +with its originating label id in ``SegmentationLabelIds``. +``WorkflowConvertVTKToUSD`` splits on that array when given ``label_names``, +which is what keeps structure identity -- and therefore per-structure +materials -- through the time series; without it a time series is split by +connectivity instead, and a component's index is no guarantee of which +structure it holds from one frame to the next. The whole-heart surface +Tutorial 4 writes beside those files is skipped: its geometry is already +covered by the per-structure ones and would be exported twice. + +Each frame is contoured from its own labelmap, so the frames agree on neither +point count nor triangulation. The stage therefore carries time-sampled +topology, which USD holds rather than interpolates, so playback snaps from +frame to frame; a surface propagated through a registration (Tutorial 1) keeps +one topology and does interpolate. + +Data Required +------------- +``tutorials/output/tutorial_04_duke_heart_labelmap/pm????/*_surfaces.vtp`` +(run ``tutorial_04_duke_heart_labelmap_to_vtk.py`` first) +""" + +# Imports +from __future__ import annotations + +import logging +from pathlib import Path +from typing import cast + +import numpy as np +import pyvista as pv + +from physiotwin4d import ( + PhysioTwin4DBase, + SegmentHeartSimplewareTrimmedBranches, + TestTools, + WorkflowConvertVTKToUSD, +) + +# Only run if this script is not imported as a module +if __name__ == "__main__": + # Data directory specification + tutorials_dir = Path(__file__).resolve().parent + + class_name = "tutorial_05_duke_heart_vtk_to_usd" + + input_dir = tutorials_dir / "output" / "tutorial_04_duke_heart_labelmap" + output_dir = tutorials_dir / "output" / "tutorial_05_duke_heart" + + log_level = logging.INFO + reporter = PhysioTwin4DBase(class_name=class_name, log_level=log_level) + + # The labelmaps were produced by Simpleware ASCardio, so its taxonomy is + # what maps each structure onto an anatomy group and a material. + segmenter = SegmentHeartSimplewareTrimmedBranches(log_level=logging.WARNING) + + output_dir.mkdir(parents=True, exist_ok=True) + + case_dirs = [input_dir / "pm0004"] + # sorted( + # path for path in input_dir.glob("pm[0-9][0-9][0-9][0-9]") if path.is_dir() + # ) + if not case_dirs: + raise FileNotFoundError( + f"No pm???? case directories found under {input_dir}.\n" + "Run tutorial_04_duke_heart_labelmap_to_vtk.py first." + ) + + # Conversion + usd_files: list[Path] = [] + last_time_codes: list[float] = [] + all_structures: set[str] = set() + label_names = segmenter.taxonomy.all_labels() + for case_dir in case_dirs: + # One file per frame, already holding every structure; sorting by name + # puts the frames in gating order. + frame_meshes = [ + cast(pv.PolyData, pv.read(str(vtp_file))) + for vtp_file in sorted(case_dir.glob("*_surfaces.vtp")) + ] + if not frame_meshes: + reporter.log_warning("%s: no per-frame surface files", case_dir.name) + continue + + # The merged files carry label ids but not names, which come from the + # taxonomy of the segmenter that produced the labelmaps. + mask_ids = { + int(label_id): label_names.get(int(label_id), f"label_{label_id}") + for mesh in frame_meshes + for label_id in np.unique(mesh.cell_data["SegmentationLabelIds"]) + } + all_structures.update(mask_ids.values()) + reporter.log_section( + f"{case_dir.name}: {len(frame_meshes)} frames, {len(mask_ids)} structures" + ) + + # One frame per time code and one cardiac cycle per second of playback. + workflow = WorkflowConvertVTKToUSD( + input_meshes=frame_meshes, + usd_project_name=case_dir.name, + output_directory=output_dir, + separate_by_connectivity=False, + appearance="anatomy", + label_names=mask_ids, + segmenter=segmenter, + frames_per_second=float(len(frame_meshes)), + log_level=log_level, + ) + results = workflow.process() + usd_files.append(Path(results["usd_file"])) + last_time_codes.append(float(len(frame_meshes) - 1)) + + reporter.log_section(f"Wrote {len(usd_files)} animated USD files to {output_dir}") + + # Testing + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + log_level=log_level, + ) + + # The first frame and the last one: a stage whose later frames were written + # against the first frame's topology renders as garbage in the second shot + # while the first still looks right. + screenshots: list[Path] = [] + if usd_files: + screenshots.append( + tt.save_screenshot_openusd( + usd_files[0], + f"{usd_files[0].stem}_usd_mesh_rendering.png", + ) + ) + screenshots.append( + tt.save_screenshot_openusd( + usd_files[0], + f"{usd_files[0].stem}_usd_mesh_rendering_last_frame.png", + time_code=last_time_codes[0], + ) + ) + + tutorial_results = { + "usd_files": usd_files, + "structures": sorted(all_structures), + "screenshots": screenshots, + } diff --git a/tutorials/tutorial_06_duke_heart_create_statistical_model.py b/tutorials/tutorial_06_duke_heart_create_statistical_model.py new file mode 100644 index 00000000..8e36b4a4 --- /dev/null +++ b/tutorials/tutorial_06_duke_heart_create_statistical_model.py @@ -0,0 +1,231 @@ +""" +Tutorial 6 (Duke Heart): Create a PCA Statistical Shape Model + +Purpose +------- +Build a PCA statistical shape model of the heart from the Duke-Heart-4DLabelmaps +population, starting from the surfaces Tutorial 4 (Duke Heart) wrote. Only each +case's reference frame is read (``*_ref_*``), so the model spans patients rather +than cardiac phases. An unbiased mean surface is built with +``WorkflowCreateMeanSurface``, and the population is decomposed into shape modes. +``ParametersDukeHeartLabelmaps.hold_out_case`` is left out of the population, so +that Tutorial 7 fitting the model to that case measures generalization rather +than reconstruction. Tutorial 7 reuses the saved ``pca_model.json``. + +The heart here is the whole heart minus its chamber cavities, the same structure +Tutorial 2 measures its distance maps to. + +Data Required +------------- +``tutorials/output/tutorial_04_duke_heart_labelmap/pm????/ +*_ref_heart_minus_interior_chambers.vtp`` +(run ``tutorial_04_duke_heart_labelmap_to_vtk.py`` first) + +Outputs (under ``tutorials/output/tutorial_06_duke_heart/``) +----------------------------------------------------------- +- ``reference_mean_surface.vtp`` - the unbiased atlas surface +- ``pca_model.json`` and ``pca_mean_surface.vtp`` - the shape model +- ``pca_mode__{minus,plus}_2sigma.vtp`` and ``pca_mode_.png`` + +Runtime +------- +An ICP and a deformable registration per case per atlas iteration, then one more +of each per case for the model. The atlas surface is cached on disk, so a +re-run only redoes the model. +""" + +# Imports +from __future__ import annotations + +import json +import logging +from pathlib import Path +from typing import Any, cast + +import numpy as np +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + ContourTools, + TestTools, + WorkflowCreateMeanSurface, + WorkflowCreateStatisticalModel, +) + +# Only run if this script is not imported as a module +if __name__ == "__main__": + # Data directory specification + repo_root = Path(__file__).resolve().parent.parent + tutorials_dir = Path(__file__).resolve().parent + + class_name = "tutorial_06_duke_heart_create_statistical_model" + + output_dir = tutorials_dir / "output" / "tutorial_06_duke_heart" + baselines_dir = repo_root / "tests" / "baselines" + + test_mode = TestTools.running_as_test() + input_dir = DUKE_HEART.input_directory(test_mode) + number_of_pca_components = DUKE_HEART.pca_components(test_mode) + + # Atlas iterations used to build the reference surface; 1 is a single + # template-biased pass. + mean_surface_iterations = 1 if test_mode else 3 + + # Points kept per surface; see the parameters module. + model_points = DUKE_HEART.model_points + + log_level = logging.INFO + + # Directory setup and data reading + + output_dir.mkdir(parents=True, exist_ok=True) + + contour_tools = ContourTools(log_level=log_level) + + # One reference-frame surface per case, less the held-out one: Tutorial 7 + # fits this model to that case, so the model must not have seen it. + sample_files = [ + path + for path in sorted(input_dir.glob("*_ref_heart_minus_interior_chambers.vtp")) + if not path.name.startswith(DUKE_HEART.hold_out_case) + ] + if test_mode: + sample_files = sample_files[:3] + if len(sample_files) < 3: + raise FileNotFoundError( + f"Need at least 3 reference-frame heart surfaces under {input_dir}; " + f"found {len(sample_files)}.\n" + "Run tutorial_04_duke_heart_labelmap_to_vtk.py first." + ) + + sample_surfaces: list[pv.DataSet] = [] + for sample_file in sample_files: + surface = cast(pv.PolyData, pv.read(str(sample_file))) + sample_surfaces.append( + contour_tools.remesh_and_smooth_surface( + surface, 1.0 - model_points / surface.n_points, 0 + ) + ) + + # The reference surface defines the topology every PCA input is expressed + # in, so picking one case makes the model inherit that case's shape. Use the + # unbiased mean of the population instead. Cached: it costs one deformable + # registration per case per atlas iteration. + reference_surface_file = output_dir / "reference_mean_surface.vtp" + if not reference_surface_file.exists(): + mean_workflow = WorkflowCreateMeanSurface( + surfaces=sample_surfaces, + template_surface=sample_surfaces[4], + log_level=log_level, + ) + mean_workflow.set_number_of_iterations(mean_surface_iterations) + mean_result = mean_workflow.process() + mean_result["mean_surface"].save(str(reference_surface_file)) + reference_surface = pv.read(str(reference_surface_file)) + + # Workflow initialization + + workflow = WorkflowCreateStatisticalModel( + sample_meshes=sample_surfaces, + reference_mesh=reference_surface, + number_of_pca_components=number_of_pca_components, + log_level=log_level, + ) + + # Workflow execution + result = workflow.process() + + # Result saving + pca_model: dict[str, Any] = result["pca_model"] + model_file = DUKE_HEART.pca_json_file + model_file.parent.mkdir(parents=True, exist_ok=True) + with model_file.open("w", encoding="utf-8") as f: + json.dump(pca_model, f, indent=2) + + mean_surface = result["pca_mean_surface"] + mean_surface_file = DUKE_HEART.pca_mean_file + mean_surface.save(str(mean_surface_file)) + + # Testing + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + baselines_dir=baselines_dir, + log_level=log_level, + ) + + screenshots: list[Path] = [] + screenshots.append( + tt.save_screenshot_mesh( + mean_surface, + "pca_mean_model.png", + camera_position="iso", + color="steelblue", + opacity=0.9, + ) + ) + + components = pca_model.get("components", []) + eigenvalues = pca_model.get("eigenvalues", []) + mean_points = np.asarray(mean_surface.points) + # PCA rank is capped by the sample count, so the model can hold fewer + # components than requested. + mode_count = min(number_of_pca_components, len(components), len(eigenvalues)) + + mode_surface_files: list[Path] = [] + xvfb_started = False + try: + pv.start_xvfb() + xvfb_started = True + except Exception: + pass + + try: + for mode_idx in range(mode_count): + sigma = float(np.sqrt(eigenvalues[mode_idx])) + mode_offsets = np.asarray(components[mode_idx]).reshape(-1, 3) + + minus_mesh = mean_surface.copy() + minus_mesh.points = mean_points - 2.0 * sigma * mode_offsets + plus_mesh = mean_surface.copy() + plus_mesh.points = mean_points + 2.0 * sigma * mode_offsets + + minus_file = output_dir / f"pca_mode_{mode_idx + 1:02d}_minus_2sigma.vtp" + plus_file = output_dir / f"pca_mode_{mode_idx + 1:02d}_plus_2sigma.vtp" + minus_mesh.save(str(minus_file)) + plus_mesh.save(str(plus_file)) + mode_surface_files.extend([minus_file, plus_file]) + + plotter = pv.Plotter(off_screen=True, window_size=[1200, 500], shape=(1, 3)) + plotter.subplot(0, 0) + plotter.add_mesh(minus_mesh, color="royalblue", opacity=0.9) + plotter.camera_position = "iso" + plotter.subplot(0, 1) + plotter.add_mesh(mean_surface, color="steelblue", opacity=0.9) + plotter.camera_position = "iso" + plotter.subplot(0, 2) + plotter.add_mesh(plus_mesh, color="coral", opacity=0.9) + plotter.camera_position = "iso" + + png_path = output_dir / f"pca_mode_{mode_idx + 1:02d}.png" + plotter.screenshot(str(png_path)) + plotter.close() + screenshots.append(png_path) + finally: + # Pair start_xvfb with cleanup, guarded like the startup above so + # environments without Xvfb (e.g. Windows, pyvista >= 0.48) are unaffected. + if xvfb_started: + try: + pv.stop_xvfb() + except Exception: + pass + + tutorial_results = { + "pca_model": pca_model, + "mean_surface": mean_surface, + "model_file": model_file, + "mean_surface_file": mean_surface_file, + "mode_surface_files": mode_surface_files, + "screenshots": screenshots, + } diff --git a/tutorials/tutorial_06_heart_create_statistical_model.py b/tutorials/tutorial_06_heart_create_statistical_model.py index 864f4c40..631a843a 100644 --- a/tutorials/tutorial_06_heart_create_statistical_model.py +++ b/tutorials/tutorial_06_heart_create_statistical_model.py @@ -4,7 +4,8 @@ Purpose ------- Build a PCA statistical shape model from a reference mesh and a small population -of sample meshes. Tutorial 7 can reuse the saved ``pca_model.json``. +of sample meshes, less ``ParametersHeartCTKCL.hold_out_case``, which Tutorial 7 +fits the model to. Tutorial 7 reuses the saved ``pca_model.json``. Data Required ------------- @@ -47,10 +48,7 @@ baselines_dir = repo_root / "tests" / "baselines" test_mode = TestTools.running_as_test() - if test_mode: - data_dir = repo_root / "data" / "test" / "KCL-Heart-Model" - else: - data_dir = repo_root / "data" / "KCL-Heart-Model" + data_dir = HEART_CT_KCL.input_directory(test_mode) number_of_pca_components = HEART_CT_KCL.pca_components(test_mode) log_level = logging.INFO @@ -73,6 +71,12 @@ sample_files = [ path for path in sample_files if path.name != reference_file.name ] + # Tutorial 7 fits this model to the held-out case, so the model must not + # have seen it. The KCL meshes carry no DIR-Lab case, so this drops nothing + # today; adding one cannot slip it in. + sample_files = [ + path for path in sample_files if HEART_CT_KCL.hold_out_case not in path.name + ] if len(sample_files) < 3: raise FileNotFoundError( f"Need at least 3 sample meshes under {sample_dir} or {data_dir}.\n" @@ -98,11 +102,12 @@ pca_model: dict[str, Any] = result["pca_model"] mean_surface: pv.PolyData = result["pca_mean_surface"] - model_file = output_dir / "pca_model.json" + model_file = HEART_CT_KCL.pca_json_file + model_file.parent.mkdir(parents=True, exist_ok=True) with model_file.open("w", encoding="utf-8") as f: json.dump(pca_model, f, indent=2) - mean_surface_file = output_dir / "pca_mean_surface.vtp" + mean_surface_file = HEART_CT_KCL.pca_mean_file mean_surface.save(str(mean_surface_file)) # Testing diff --git a/tutorials/tutorial_06_lung_create_statistical_model.py b/tutorials/tutorial_06_lung_create_statistical_model.py index c2d8ec45..5021ac22 100644 --- a/tutorials/tutorial_06_lung_create_statistical_model.py +++ b/tutorials/tutorial_06_lung_create_statistical_model.py @@ -3,7 +3,8 @@ Purpose ------- -Build a PCA statistical shape model of the lungs from the DIR-Lab population. +Build a PCA statistical shape model of the lungs from the DIR-Lab population, +less ``ParametersLungCTDirLab.hold_out_case``, which Tutorial 7 fits it to. Each case's T70 phase is segmented, an unbiased mean surface is built with ``WorkflowCreateMeanSurface``, and the population is decomposed into shape modes. Tutorials 7 and 8 reuse the saved ``pca_model.json``. @@ -67,11 +68,10 @@ output_dir = tutorials_dir / "output" / "tutorial_06_lung" baselines_dir = repo_root / "tests" / "baselines" - data_dir = repo_root / "data" / "DirLab-4DCT" + test_mode = TestTools.running_as_test() + data_dir = LUNG_CT_DIRLAB.input_directory(test_mode) - number_of_pca_components = LUNG_CT_DIRLAB.pca_components( - TestTools.running_as_test() - ) + number_of_pca_components = LUNG_CT_DIRLAB.pca_components(test_mode) # Atlas iterations used to build the reference surface; 1 is a single # template-biased pass. @@ -91,7 +91,14 @@ contour_tools = ContourTools(log_level=log_level) - sample_image_files = sorted(data_dir.glob("Case*T70.mha")) + # Tutorial 7 fits this model to the held-out study, so the model must not + # have seen it. That study lives in another dataset, so this drops nothing + # today; moving it in here cannot slip it in. + sample_image_files = [ + path + for path in sorted(data_dir.glob("Case*T70.mha")) + if path.name != LUNG_CT_DIRLAB.hold_out_case + ] sample_surfaces = [] for sample_image_file in sample_image_files: sample_surface_file = output_dir / f"{sample_image_file.stem}.vtp" @@ -140,12 +147,13 @@ # Result saving pca_model: dict[str, Any] = result["pca_model"] - model_file = output_dir / "pca_model.json" + model_file = LUNG_CT_DIRLAB.pca_json_file + model_file.parent.mkdir(parents=True, exist_ok=True) with model_file.open("w", encoding="utf-8") as f: json.dump(pca_model, f, indent=2) mean_surface = result["pca_mean_surface"] - mean_surface_file = output_dir / "pca_mean_surface.vtp" + mean_surface_file = LUNG_CT_DIRLAB.pca_mean_file mean_surface.save(str(mean_surface_file)) # Testing diff --git a/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py new file mode 100644 index 00000000..37c89da2 --- /dev/null +++ b/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py @@ -0,0 +1,250 @@ +""" +Tutorial 7 (Duke Heart): Fit Statistical Shape Model to Patient Data + +Purpose +------- +Duke counterpart of ``tutorial_07_heart_fit_statistical_model_to_patient.py``. +The heart PCA model built by ``tutorial_06_duke_heart_create_statistical_model.py`` +is fitted to one Duke-Heart-4DLabelmaps case, whose reference frame plays the +patient. + +Duke-Heart-4DLabelmaps ships labelmaps rather than CT, so there is nothing to +segment: the patient surface is contoured straight from the case's labelmap, +with the same labels dropped that Tutorial 4 drops, so the patient and the model +describe the same structure. Nothing being segmented also means there +is no patient intensity image; the workflow then rasterizes its own reference +grid from the patient surface. + +The patient is ``ParametersDukeHeartLabelmaps.hold_out_case``, which Tutorial 6 +leaves out of the population it builds the model from, so this fit measures +generalization rather than reconstruction. + +Data Required +------------- +PCA model: Tutorial 6 output (``output/tutorial_06_duke_heart/pca_model.json``, +``pca_mean_surface.vtp``) +Patient: ``data/Duke-Heart-4DLabelmaps//*_ref_labelmap.nii.gz`` +ICON weights: ``tutorial_02_duke_heart_distancemap_finetune_icon.py`` output +(``network_weights/icon_duke_heart_distancemap/ +icon_duke_heart_distancemap_model/checkpoints/network_weights_final.trch``), +optional -- the stock uniGradICON weights are used when it is absent. +""" + +# Imports +from __future__ import annotations + +import json +import logging +from pathlib import Path +from typing import Any, Optional, cast + +import itk +import numpy as np +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + ContourTools, + TestTools, + WorkflowFitStatisticalModelToPatient, +) + +# Only run if this script is not imported as a module +if __name__ == "__main__": + # Data directory specification + repo_root = Path(__file__).resolve().parent.parent + tutorials_dir = Path(__file__).resolve().parent + + project_name = "tutorial_07_duke_heart" + + output_dir = tutorials_dir / "output" / project_name + baselines_dir = repo_root / "tests" / "baselines" + + # PCA model + mean surface produced by Tutorial 6 (Duke Heart). + pca_json = DUKE_HEART.pca_json_file + pca_mean_file = DUKE_HEART.pca_mean_file + + # The case whose reference frame plays the patient: the one Tutorial 6 + # leaves out of the model, so this fit is out of sample. + patient_case = DUKE_HEART.hold_out_case + + test_mode = TestTools.running_as_test() + number_of_pca_components = DUKE_HEART.pca_components(test_mode) + data_dir = DUKE_HEART.hold_out_directory(test_mode) + + # Labels left out of the whole-heart structure, the same ones Tutorial 4 + # drops when it builds the surfaces the model was trained on, so the patient + # and the model describe the same structure. + interior_object_ids = DUKE_HEART.interior_object_ids + + # Contouring grid, shared with Tutorial 4 so the patient surface carries the + # same level of detail as the model's training surfaces. + surface_spacing_mm = DUKE_HEART.surface_spacing_mm + smoothing_iterations = DUKE_HEART.surface_smoothing_iterations + + # Distance-map weights finetuned by + # tutorial_02_duke_heart_distancemap_finetune_icon.py. The heart has its + # own finetuning run rather than reusing the lung one's: the heart + # registration mask is far tighter, so heart distance maps saturate over a + # shorter radius and do not share an intensity distribution with lung ones. + icon_weights_path = ( + tutorials_dir + / "network_weights" + / "icon_duke_heart_distancemap" + / "icon_duke_heart_distancemap_model" + / "checkpoints" + / "network_weights_final.trch" + ) + + log_level = logging.INFO + + contour_tools = ContourTools(log_level=log_level) + + # Directory setup and data reading + + output_dir.mkdir(parents=True, exist_ok=True) + + if not pca_mean_file.exists(): + raise FileNotFoundError( + f"Tutorial 6 PCA mean surface not found: {pca_mean_file}\n" + "Run tutorials/tutorial_06_duke_heart_create_statistical_model.py first." + ) + pca_mean = cast(pv.DataSet, pv.read(str(pca_mean_file))) + + pca_model: Optional[dict[str, Any]] = None + if pca_json.exists(): + with pca_json.open(encoding="utf-8") as f: + pca_model = json.load(f) + + patient_labelmap_files = sorted( + (data_dir / patient_case).glob("*_ref_labelmap.nii.gz") + ) + if not patient_labelmap_files: + raise FileNotFoundError( + f"No reference-frame labelmap for {patient_case} under {data_dir}.\n" + "See data/Duke-Heart-4DLabelmaps/README.md." + ) + patient_labelmap = itk.imread(str(patient_labelmap_files[0])) + + # The patient surface, contoured once and cached: the whole heart minus its + # chamber cavities, which is the structure the model describes. + heart_surface_file = output_dir / f"{project_name}_heart_surface.vtp" + if not heart_surface_file.exists(): + labels = itk.GetArrayViewFromImage(patient_labelmap) + heart_ids = [ + int(value) + for value in np.unique(labels) + if value != 0 and int(value) not in interior_object_ids + ] + heart_mask = itk.GetImageFromArray(np.isin(labels, heart_ids).astype(np.uint8)) + heart_mask.CopyInformation(patient_labelmap) + heart_surface = contour_tools.extract_label_surfaces( + heart_mask, + isotropic_spacing_mm=surface_spacing_mm, + smoothing_iterations=smoothing_iterations, + )[1] + heart_surface.save(str(heart_surface_file)) + heart_surface = cast(pv.PolyData, pv.read(str(heart_surface_file))) + + # Workflow initialization + + workflow = WorkflowFitStatisticalModelToPatient( + template_model=pca_mean, + patient_models=[heart_surface], + # This dataset carries no intensity image, so the workflow rasterizes + # its own reference grid from the patient surface. + patient_image=None, + patient_labelmap=patient_labelmap, + log_level=log_level, + # The labels the whole-heart surface leaves out are the ones a distance + # map must not measure to either. + labelmap_interior_object_ids=interior_object_ids, + ) + workflow.set_mask_dilation_mm(DUKE_HEART.mask_dilation_mm) + workflow.set_distancemap_squared_max(DUKE_HEART.distancemap_squared_max) + if pca_model is not None: + workflow.set_use_pca_registration( + use_pca_registration=True, + pca_model=pca_model, + number_of_pca_components=number_of_pca_components, + use_surface=False, + ) + + # The labelmap-to-labelmap stage registers distance maps, not intensities, + # so it uses the distance-map-finetuned weights when they exist; without + # them the tutorial still runs, on the stock uniGradICON weights. + if icon_weights_path.exists(): + workflow.set_labelmap_to_labelmap_icon_weights_path(str(icon_weights_path)) + else: + workflow.log_warning( + "Finetuned distance-map ICON weights not found at %s; fitting with " + "the stock uniGradICON weights. Run " + "tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py to " + "create them.", + icon_weights_path, + ) + + # Workflow execution + workflow_results = workflow.process() + + # Result saving + registered_coefficients = workflow.pca_coefficients + if registered_coefficients is not None: + registered_coefficients_path = ( + output_dir / f"{project_name}_registered_coefficients.json" + ) + with registered_coefficients_path.open(mode="w", encoding="utf-8") as f: + json.dump(registered_coefficients.tolist(), f) + + # The Duke heart PCA model from Tutorial 6 is built from surfaces only, so + # the model *is* a surface: the volume mesh and its bounding surface are the + # same geometry and only the .vtp surfaces are written. + template_surface = workflow.pca_template_model_surface + assert template_surface is not None, ( + "pca_template_model_surface must be set after process()" + ) + template_surface.save(str(output_dir / f"{project_name}_template_surface.vtp")) + + registered_surface = workflow_results["registered_template_model_surface"] + registered_surface.save( + str(output_dir / f"{project_name}_template_surface_registered.vtp") + ) + + # Testing + TestTools( + class_name=project_name, + results_dir=output_dir, + baselines_dir=baselines_dir, + log_level=log_level, + ) + + try: + pv.start_xvfb() + except Exception: + pass + + screenshots: list[Path] = [] + + before_path = output_dir / f"{project_name}_model_before_registration.png" + plotter = pv.Plotter(off_screen=True, window_size=[800, 600]) + plotter.add_mesh(pca_mean, color="dodgerblue", opacity=0.6) + plotter.add_mesh(heart_surface, color="tomato", opacity=0.6) + plotter.camera_position = "iso" + plotter.screenshot(str(before_path)) + plotter.close() + screenshots.append(before_path) + + after_path = output_dir / f"{project_name}_model_after_registration.png" + plotter = pv.Plotter(off_screen=True, window_size=[800, 600]) + plotter.add_mesh(registered_surface, color="limegreen", opacity=0.7) + plotter.add_mesh(heart_surface, color="tomato", opacity=0.4) + plotter.camera_position = "iso" + plotter.screenshot(str(after_path)) + plotter.close() + screenshots.append(after_path) + + tutorial_results = { + "registered_surface": registered_surface, + "registered_coefficients": registered_coefficients, + "screenshots": screenshots, + } diff --git a/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py index 8e446a92..bda6901e 100644 --- a/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py +++ b/tutorials/tutorial_07_heart_fit_statistical_model_to_patient.py @@ -52,17 +52,13 @@ baselines_dir = repo_root / "tests" / "baselines" # PCA model + mean surface produced by Tutorial 6. - pca_json = tutorials_dir / "output" / "tutorial_06_heart" / "pca_model.json" - pca_mean_file = ( - tutorials_dir / "output" / "tutorial_06_heart" / "pca_mean_surface.vtp" - ) + pca_json = HEART_CT_KCL.pca_json_file + pca_mean_file = HEART_CT_KCL.pca_mean_file test_mode = TestTools.running_as_test() - if test_mode: - data_dir = repo_root / "data" / "test" / "DirLab-4DCT" - else: - data_dir = repo_root / "data" / "DirLab-4DCT" - patient_image_file = data_dir / "Case1Pack_T70.mha" + data_dir = HEART_CT_KCL.hold_out_directory(test_mode) + # The case Tutorial 6 leaves out of the model, so this fit is out of sample. + patient_image_file = data_dir / f"{HEART_CT_KCL.hold_out_case}_T70.mha" # Distance-map weights finetuned by # tutorial_02_duke_heart_distancemap_finetune_icon.py. The heart has its own @@ -133,7 +129,10 @@ ) contour_tools = ContourTools() - heart_surface = contour_tools.extract_contours(labelmap_image=heart_labelmap) + heart_surface = contour_tools.extract_contours( + labelmap_image=heart_labelmap, + surface_reduction_rate=HEART_CT_KCL.surface_reduction_rate, + ) heart_surface.save(output_dir / f"{project_name}_heart_surface.vtp") else: @@ -165,6 +164,7 @@ workflow.set_use_pca_registration( use_pca_registration=True, pca_model=pca_model, + number_of_pca_components=HEART_CT_KCL.pca_components(test_mode), use_surface=False, ) diff --git a/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py index 5a33a947..f8f048e0 100644 --- a/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py +++ b/tutorials/tutorial_07_lung_fit_statistical_model_to_patient.py @@ -58,15 +58,16 @@ baselines_dir = repo_root / "tests" / "baselines" # PCA model + mean surface produced by Tutorial 6 (lung). - tutorial_06_dir = tutorials_dir / "output" / "tutorial_06_lung" - pca_json = tutorial_06_dir / "pca_model.json" - pca_mean_file = tutorial_06_dir / "pca_mean_surface.vtp" + pca_json = LUNG_CT_DIRLAB.pca_json_file + pca_mean_file = LUNG_CT_DIRLAB.pca_mean_file - number_of_pca_components = LUNG_CT_DIRLAB.pca_components( - TestTools.running_as_test() - ) + test_mode = TestTools.running_as_test() + number_of_pca_components = LUNG_CT_DIRLAB.pca_components(test_mode) - patient_image_file = repo_root / "data" / "Chest-CT" / "Chest-CT.mha" + # The study Tutorial 6 leaves out of the model, so this fit is out of sample. + patient_image_file = ( + LUNG_CT_DIRLAB.hold_out_directory(test_mode) / LUNG_CT_DIRLAB.hold_out_case + ) # Distance-map weights finetuned on DIR-Lab by # tutorial_02_lung_distancemap_finetune_icon.py; see @@ -119,6 +120,7 @@ segmentation_result = segmentation_workflow.process( input_image=patient_image, anatomy_groups=["lung"], + surface_reduction_rate=LUNG_CT_DIRLAB.surface_reduction_rate, extract_label_surfaces=True, ) contour_tools.save_combined_surfaces( diff --git a/tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py b/tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py new file mode 100644 index 00000000..568a684c --- /dev/null +++ b/tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py @@ -0,0 +1,362 @@ +""" +Tutorial 8 (Duke Heart): Fit the Heart SSM and Propagate It Through Cardiac Phases + +Purpose +------- +Duke counterpart of ``tutorial_08_lung_fit_model_to_4d_patients.py``, run on the +Duke-Heart-4DLabelmaps cohort. For each case it produces one +statistical-shape-model (SSM) surface per gated cardiac frame: + +1. Fit the heart PCA model to the case's reference frame + (``*_ref_labelmap.nii.gz``). The whole heart minus its chamber cavities is + contoured from that frame and the PCA model built by + ``tutorial_06_duke_heart_create_statistical_model.py`` is fitted to it with + PCA-based registration (``WorkflowFitStatisticalModelToPatient`` with + ``use_pca_registration=True``). This yields the case's PCA coefficients plus + the fitted SSM surface, sharing the model's fixed topology. + +2. Propagate that surface to every gated frame. This data ships labelmaps + rather than CT, so there is no intensity image to register: each frame's + heart surface is contoured the same way and the fitted SSM surface is + registered to it with ``RegisterModelsDistanceMaps`` (Greedy affine, then + ICON on the distance maps), which warps the SSM while keeping its topology. + The reference frame keeps the fitted surface itself rather than being + registered to its own contour. + +Every frame's distance maps are rasterized on one grid built around the +reference frame's heart, so the phases of a case are registered in a common +space even though their labelmaps carry different slice pitches. + +Statistical models are typically dense tetrahedral volume meshes, written as +``.vtu``. The heart PCA model from Tutorial 6 (Duke Heart) is built from +surfaces only, so every model here is a surface and every model file is a +``.vtp``. + +Data Required +------------- +data: ``data/Duke-Heart-4DLabelmaps/pm????/*_labelmap.nii.gz`` +PCA model: Tutorial 6 (Duke Heart) output +(``output/tutorial_06_duke_heart/pca_model.json``, ``pca_mean_surface.vtp``) +Surfaces: Tutorial 4 (Duke Heart) ``outputs = "full"`` output, optional -- any +frame it did not contour is contoured here instead. +ICON weights: Tutorial 2 (Duke Heart) output +(``network_weights/icon_duke_heart_distancemap/ +icon_duke_heart_distancemap_model/checkpoints/network_weights_final.trch``), +optional -- the stock uniGradICON weights are used when it is absent. + +Outputs (per case, under ``output/tutorial_08_duke_heart//``) +------------------------------------------------------------------ + * ``_ssm_pca_coefficients.json`` - fitted PCA coefficient vector + * ``_ssm_pca_surface.vtp`` - PCA template before the final warp + * ``_ssm_surface.vtp`` - fitted reference-frame SSM surface + * ``_ssm_surface.vtp`` - SSM warped to that gated frame + * ``_heart_surface.vtp`` - contoured frame surface, cached +""" + +# Imports +from __future__ import annotations + +import json +import logging +from pathlib import Path +from typing import Any, cast + +import itk +import numpy as np +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + ContourTools, + RegisterModelsDistanceMaps, + TestTools, + WorkflowFitStatisticalModelToPatient, +) + +# Structure name Tutorial 4 (Duke Heart) writes its whole-heart surfaces under. +WHOLE_HEART_NAME = "heart_minus_interior_chambers" +LABELMAP_SUFFIX = "_labelmap.nii.gz" + +# Only run if this script is not imported as a module + +# The registration backends spawn worker processes. On Windows the spawn start +# method re-imports this script in each child; without the +# __name__ == "__main__" guard around top-level work, that re-import would +# restart the whole cohort in every worker. +if __name__ == "__main__": + # Data directory specification + repo_root = Path(__file__).resolve().parent.parent + tutorials_dir = Path(__file__).resolve().parent + + class_name = "tutorial_08_duke_heart_fit_model_to_4d_patients" + + output_dir = tutorials_dir / "output" / "tutorial_08_duke_heart" + baselines_dir = repo_root / "tests" / "baselines" + + test_mode = TestTools.running_as_test() + # The gated labelmaps, one directory per case. + data_dir = DUKE_HEART.hold_out_directory(test_mode) + # Tutorial 4's surfaces, read when its "full" pass wrote the frame. + tutorial_04_dir = DUKE_HEART.input_directory(test_mode) + + # PCA model + mean surface produced by Tutorial 6 (Duke Heart). + pca_model_file = DUKE_HEART.pca_json_file + pca_mean_file = DUKE_HEART.pca_mean_file + + number_of_pca_components = DUKE_HEART.pca_components(test_mode) + + # Labels left out of the whole-heart structure, the same ones Tutorials 4, + # 6 and 7 drop, so the frames and the model describe the same structure. + interior_object_ids = DUKE_HEART.interior_object_ids + + # Contouring grid, shared with Tutorial 4 so every surface here carries the + # same level of detail as the model's training surfaces. + surface_spacing_mm = DUKE_HEART.surface_spacing_mm + smoothing_iterations = DUKE_HEART.surface_smoothing_iterations + + # Pitch of the grid the phase distance maps are rasterized on. Coarser than + # the contouring pitch: it carries a distance field, not a boundary. + registration_spacing_mm = 1.0 + + # Distance-map weights finetuned by + # tutorial_02_duke_heart_distancemap_finetune_icon.py, used both by the + # labelmap-to-labelmap stage of the SSM fit and by the phase registrations. + icon_weights_path = ( + tutorials_dir + / "network_weights" + / "icon_duke_heart_distancemap" + / "icon_duke_heart_distancemap_model" + / "checkpoints" + / "network_weights_final.trch" + ) + + log_level = logging.INFO + + logging.basicConfig(level=log_level) + logger = logging.getLogger(class_name) + + # Directory setup and data reading + + output_dir.mkdir(parents=True, exist_ok=True) + + for required_file in (pca_model_file, pca_mean_file): + if not required_file.exists(): + raise FileNotFoundError( + f"Tutorial 6 output not found: {required_file}\n" + "Run tutorials/tutorial_06_duke_heart_create_statistical_model.py " + "first." + ) + pca_mean_surface = cast(pv.DataSet, pv.read(str(pca_mean_file))) + with pca_model_file.open(encoding="utf-8") as f: + pca_model = json.load(f) + + # The Tutorial 2 distance-map weights are used when they exist; without them + # the tutorial still runs, on the stock uniGradICON weights. + use_finetuned_weights = icon_weights_path.exists() + if use_finetuned_weights: + logger.info( + "Registering with finetuned distance-map ICON weights: %s", + icon_weights_path, + ) + else: + logger.warning( + "Finetuned distance-map ICON weights not found at %s; registering " + "with the stock uniGradICON weights. Run " + "tutorials/tutorial_02_duke_heart_distancemap_finetune_icon.py to " + "create them.", + icon_weights_path, + ) + + case_dirs = sorted( + path for path in data_dir.glob("pm[0-9][0-9][0-9][0-9]") if path.is_dir() + ) + if not case_dirs: + raise FileNotFoundError( + f"No pm???? case directories found under {data_dir}.\n" + "See data/Duke-Heart-4DLabelmaps/README.md." + ) + + contour_tools = ContourTools(log_level=log_level) + + def heart_surface_for(labelmap_file: Path, case_output_dir: Path) -> pv.PolyData: + """Return one frame's whole heart, minus its chamber cavities. + + Tutorial 4's ``"full"`` pass contours this surface for every gated + frame, so its output is read when present. Otherwise the surface is + contoured here and cached beside this tutorial's other outputs, so a + re-run pays for it once. + """ + stem = labelmap_file.name[: -len(LABELMAP_SUFFIX)] + tutorial_04_file = tutorial_04_dir / f"{stem}_{WHOLE_HEART_NAME}.vtp" + if tutorial_04_file.exists(): + return cast(pv.PolyData, pv.read(str(tutorial_04_file))) + + surface_file = case_output_dir / f"{stem}_heart_surface.vtp" + if not surface_file.exists(): + labelmap = itk.imread(str(labelmap_file)) + labels = itk.GetArrayViewFromImage(labelmap) + heart_ids = [ + int(value) + for value in np.unique(labels) + if value != 0 and int(value) not in interior_object_ids + ] + heart_mask = itk.GetImageFromArray( + np.isin(labels, heart_ids).astype(np.uint8) + ) + heart_mask.CopyInformation(labelmap) + heart_surface = contour_tools.extract_label_surfaces( + heart_mask, + isotropic_spacing_mm=surface_spacing_mm, + smoothing_iterations=smoothing_iterations, + )[1] + heart_surface.save(str(surface_file)) + return cast(pv.PolyData, pv.read(str(surface_file))) + + tutorial_results: dict[str, Any] = {"cases": {}, "screenshots": []} + + for case_dir in case_dirs: + case_id = case_dir.name + frame_files = sorted(case_dir.glob(f"*{LABELMAP_SUFFIX}")) + reference_files = [ + path for path in frame_files if path.name.endswith(f"_ref{LABELMAP_SUFFIX}") + ] + if not reference_files: + logger.warning("Skipping %s: no *_ref_labelmap.nii.gz frame", case_id) + continue + reference_file = reference_files[0] + + logger.info("%s", "=" * 48) + logger.info("Processing case %s: %d gated frames", case_id, len(frame_files)) + logger.info("%s", "=" * 48) + + case_output_dir = output_dir / case_id + case_output_dir.mkdir(parents=True, exist_ok=True) + + # Step 1: contour the reference frame's heart. + reference_labelmap = itk.imread(str(reference_file)) + reference_surface = heart_surface_for(reference_file, case_output_dir) + + # Step 2: fit the statistical model to the reference frame. This data + # carries no intensity image, so the workflow rasterizes its own + # reference grid from the patient surface. + fit_workflow = WorkflowFitStatisticalModelToPatient( + template_model=pca_mean_surface, + patient_models=[reference_surface], + patient_image=None, + patient_labelmap=reference_labelmap, + labelmap_interior_object_ids=interior_object_ids, + log_level=log_level, + ) + fit_workflow.set_use_pca_registration( + use_pca_registration=True, + pca_model=pca_model, + number_of_pca_components=number_of_pca_components, + use_surface=False, + ) + fit_workflow.set_mask_dilation_mm(DUKE_HEART.mask_dilation_mm) + fit_workflow.set_distancemap_squared_max(DUKE_HEART.distancemap_squared_max) + if use_finetuned_weights: + fit_workflow.set_labelmap_to_labelmap_icon_weights_path( + str(icon_weights_path) + ) + fit_result = fit_workflow.process() + + pca_coefficients_file = case_output_dir / f"{case_id}_ssm_pca_coefficients.json" + pca_coefficients = fit_workflow.pca_coefficients + assert pca_coefficients is not None + with pca_coefficients_file.open(mode="w", encoding="utf-8") as f: + json.dump(pca_coefficients.tolist(), f) + + pca_template_surface = fit_workflow.pca_template_model_surface + assert pca_template_surface is not None + pca_template_surface.save( + str(case_output_dir / f"{case_id}_ssm_pca_surface.vtp") + ) + + # Typically the SSM is a dense tetrahedral volume mesh, saved as .vtu, + # and its bounding surface is saved separately as .vtp. The heart PCA + # model from Tutorial 6 (Duke Heart) is built from surfaces only, so + # here the model *is* a surface: "registered_template_model" and + # "registered_template_model_surface" are the same geometry, and only + # the .vtp surface is written. + ssm_surface_fitted = fit_result["registered_template_model_surface"] + ssm_surface_file = case_output_dir / f"{case_id}_ssm_surface.vtp" + ssm_surface_fitted.save(str(ssm_surface_file)) + + # Step 3: warp the fitted SSM surface onto every gated frame. One grid + # is built around the reference frame's heart and reused by every frame, + # so the whole case is registered in a common space; its buffer holds + # the frames the heart moves into. + registration_grid = contour_tools.create_reference_image( + mesh=reference_surface, + spatial_resolution=registration_spacing_mm, + buffer_factor=0.25, + ptype=itk.F, + ) + + phase_outputs = [] + for frame_file in frame_files: + stem = frame_file.name[: -len(LABELMAP_SUFFIX)] + if frame_file == reference_file: + # The fit already placed the SSM on this frame. + logger.info("Case %s: reference frame %s", case_id, stem) + phase_surface = ssm_surface_fitted + else: + logger.info("Case %s: warping to frame %s", case_id, stem) + registrar = RegisterModelsDistanceMaps( + moving_model=ssm_surface_fitted, + fixed_model=heart_surface_for(frame_file, case_output_dir), + reference_image=registration_grid, + distance_squared_max=DUKE_HEART.distancemap_squared_max, + mask_dilation_mm=DUKE_HEART.mask_dilation_mm, + log_level=log_level, + ) + if use_finetuned_weights: + registrar.set_icon_weights_path(str(icon_weights_path)) + phase_surface = registrar.register(transform_type="Deformable")[ + "registered_model" + ] + + surface_file = case_output_dir / f"{stem}_ssm_surface.vtp" + phase_surface.save(str(surface_file)) + phase_outputs.append({"frame_stem": stem, "surface_file": surface_file}) + + tutorial_results["cases"][case_id] = { + "pca_coefficients_file": pca_coefficients_file, + "ssm_surface_file": ssm_surface_file, + "phase_outputs": phase_outputs, + } + + if not tutorial_results["cases"]: + raise RuntimeError( + f"No case under {data_dir} carried a reference frame; nothing was fitted." + ) + + # Testing + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + baselines_dir=baselines_dir, + log_level=log_level, + ) + + last_case = list(tutorial_results["cases"].values())[-1] + tutorial_results["screenshots"] = [ + tt.save_screenshot_mesh( + cast(pv.DataSet, pv.read(str(last_case["ssm_surface_file"]))), + "ssm_surface_reference.png", + camera_position="iso", + color="steelblue", + opacity=0.9, + ), + tt.save_screenshot_mesh( + cast( + pv.DataSet, + pv.read(str(last_case["phase_outputs"][0]["surface_file"])), + ), + "ssm_surface_first_phase.png", + camera_position="iso", + color="limegreen", + opacity=0.9, + ), + ] diff --git a/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py b/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py index 77af9dad..8c79268a 100644 --- a/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py +++ b/tutorials/tutorial_08_lung_fit_model_to_4d_patients.py @@ -87,9 +87,10 @@ baselines_dir = repo_root / "tests" / "baselines" # PCA model + mean surface produced by Tutorial 6 (lung). - tutorial_06_dir = tutorials_dir / "output" / "tutorial_06_lung" - pca_model_file = tutorial_06_dir / "pca_model.json" - pca_mean_file = tutorial_06_dir / "pca_mean_surface.vtp" + pca_model_file = LUNG_CT_DIRLAB.pca_json_file + pca_mean_file = LUNG_CT_DIRLAB.pca_mean_file + # Tutorial 6 caches one segmentation per case beside its model. + tutorial_06_dir = pca_model_file.parent # Distance-map weights finetuned on DIR-Lab by # tutorial_02_lung_distancemap_finetune_icon.py, used by the @@ -186,6 +187,7 @@ segmentation_result = segmentation_workflow.process( input_image=reference_image, anatomy_groups=["lung"], + surface_reduction_rate=LUNG_CT_DIRLAB.surface_reduction_rate, extract_label_surfaces=True, ) contour_tools.save_combined_surfaces( diff --git a/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py b/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py new file mode 100644 index 00000000..e319bc79 --- /dev/null +++ b/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py @@ -0,0 +1,337 @@ +""" +Tutorial 9 (Duke Heart, MGN): Train a PhysicsNeMo MeshGraphNet on the Fitted SSM + +Purpose +------- +Duke counterpart of ``tutorial_09_lung_train_physicsnemo_mgn.py``, run on the +Duke-Heart-4DLabelmaps cohort. A thin driver over the reusable +:class:`physiotwin4d.WorkflowTrainPhysicsNeMo` workflow: + +1. Discover the per-frame SSM surfaces produced by Tutorial 8 (Duke Heart) + (``tutorial_08_duke_heart_fit_model_to_4d_patients.py``), write the training + target for each frame, and write one JSON manifest per case. The target here + is the per-vertex displacement from the case's reference surface, stored as a + ``displacement`` point-data array -- the workflow reads targets verbatim and + never derives them. Cardiac stages are parsed from the ``g{PPP}`` gate tag + of the labelmap filenames and written explicitly into the manifest (the + workflow never parses filenames). + +2. Split the cases into train and held-out test -- plus an optional validation + set, empty by default, which is what makes the intermittent validation RMSE + read ``n/a`` -- and train the MeshGraphNet (``WorkflowTrainPhysicsNeMo`` + driving ``TrainPhysicsNeMoMGN``). + +3. Evaluate the held-out test cases against their ground-truth frames with + :class:`physiotwin4d.WorkflowInferPhysicsNeMo` wrapped in + :class:`physiotwin4d.WorkflowInferMovement`. + +Why a GNN? +---------- +The SSM surface has a fixed topology across all cases and the myocardium is a +continuum: adjacent vertices co-vary smoothly. MeshGraphNet encodes that prior +directly by passing messages along mesh edges, giving an explicit +continuum-deformation inductive bias the MLP must infer from coordinates alone. + +Node features (per vertex): [mean_shape_x, mean_shape_y, mean_shape_z, pca_c1 ... pca_cN, stage] +Edge features (per edge): [rel_x, rel_y, rel_z, distance] (from the mean shape) +Output (per vertex): [dx, dy, dz] (displacement in mm) + +Extra Install Required +---------------------- +PhysicsNeMo and PyTorch Geometric must be installed:: + + pip install "physiotwin4d[physicsnemo]" + +Data Required +------------- +SSM surfaces: Tutorial 8 (Duke Heart) output +(``output/tutorial_08_duke_heart/pm????/``) +PCA mean surface: Tutorial 6 (Duke Heart) output +(``output/tutorial_06_duke_heart/pca_mean_surface.vtp``, alongside +``pca_model.json``) + +Outputs +------- +Manifests and per-frame targets are written under +``output/tutorial_09_duke_heart_mgn/manifests_mgn/``: + + * ``pm????_manifest.json`` - per-case training manifest + * ``_ssm_surface_target.vtp`` - per-frame displacement targets + +The evaluation of the held-out cases lands in +``output/tutorial_09_duke_heart_mgn/eval_mgn/pm????/``. + +The model itself is written to ``ParametersDukeHeartLabelmaps.mgn_weights_dir`` +(``network_weights/physicsnemo_mgn_duke_heart_motion/``), or to a fresh ``..._1`` +sibling when resuming (see ``resume_from``), which is what ``tutorial_results`` +reports as ``model_directory``: + + * ``mgn_stage_model.pt`` - trained MeshGraphNet checkpoint + * ``pca_mean_surface.vtp``, ``pca_model.json`` - PCA assets copied beside it, + so the directory is self-contained for inference +""" + +# Imports +from __future__ import annotations + +import json +import logging +from pathlib import Path +from typing import Any, Optional, cast + +import numpy as np +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + TestTools, + TrainPhysicsNeMoMGN, + WorkflowInferMovement, + WorkflowInferPhysicsNeMo, + WorkflowTrainPhysicsNeMo, +) + +# Point-data array the tutorial writes its targets into and the manifests name. +TARGET_ARRAY = "displacement" + +# Gated frames carry a ``g{PPP}`` tag naming their percentage of the R-R +# interval; this is what a per-frame SSM surface is matched and staged by. +PHASE_SURFACE_PATTERN = "*_g[0-9][0-9][0-9]_*_ssm_surface.vtp" + + +def _cardiac_stage_from_filename(surface_file: Path) -> float: + """Extract the normalized cardiac stage [0, 1] from a ``g{PPP}`` filename stem.""" + for part in surface_file.stem.split("_"): + if part.startswith("g") and part[1:].isdigit(): + return int(part[1:]) / 100.0 + raise ValueError(f"Cannot parse cardiac gate from filename: {surface_file}") + + +def _write_target_mesh( + phase_file: Path, ref_points: np.ndarray, targets_dir: Path +) -> Path: + """Write one frame's training target and return the mesh path. + + The target is the per-vertex displacement from the case's reference surface, + stored as the ``TARGET_ARRAY`` point-data array on a copy of the frame + surface. Any other per-vertex quantity could be written here instead -- the + training workflow reads whatever array the manifest names. + """ + phase_mesh = pv.read(str(phase_file)) + phase_points = np.asarray(phase_mesh.points, dtype=np.float32) + phase_mesh.point_data[TARGET_ARRAY] = phase_points - ref_points + target_path = targets_dir / f"{phase_file.stem}_target.vtp" + phase_mesh.save(str(target_path)) + return target_path + + +def _write_case_manifest( + case_dir: Path, manifests_dir: Path, logger: logging.Logger +) -> Optional[Path]: + """Write a per-case manifest JSON; return its path (or None if incomplete). + + A case needs a reference SSM surface, a PCA coefficient file, and at least + two gated-frame surfaces. A case that is missing any of them is skipped + with the reason logged, so a half-finished Tutorial 8 run is distinguishable + from one that never ran. + """ + case_id = case_dir.name + ref_file = case_dir / f"{case_id}_ssm_surface.vtp" + pca_file = case_dir / f"{case_id}_ssm_pca_coefficients.json" + phase_files = sorted(case_dir.glob(PHASE_SURFACE_PATTERN)) + missing = [] + if not ref_file.exists(): + missing.append(f"reference surface {ref_file.name}") + if not pca_file.exists(): + missing.append(f"PCA coefficients {pca_file.name}") + if len(phase_files) < 2: + missing.append(f"at least 2 frame surfaces (found {len(phase_files)})") + if missing: + logger.warning("Skipping %s: missing %s", case_id, "; ".join(missing)) + return None + + manifests_dir.mkdir(parents=True, exist_ok=True) + ref_points = np.asarray(pv.read(str(ref_file)).points, dtype=np.float32) + manifest = { + "subject_id": case_id, + "reference_mesh": str(ref_file), + "pca_coefficients": str(pca_file), + "target_array": TARGET_ARRAY, + "phases": [ + { + "mesh": str(_write_target_mesh(phase_file, ref_points, manifests_dir)), + "stage": _cardiac_stage_from_filename(phase_file), + } + for phase_file in phase_files + ], + } + manifest_path = manifests_dir / f"{case_id}_manifest.json" + manifest_path.write_text(json.dumps(manifest, indent=2), encoding="utf-8") + return manifest_path + + +# Only run if this script is not imported as a module + +# PhysicsNeMo and torch spawn worker processes for data loading. On Windows the +# spawn start method re-imports this script in each child; without the +# __name__ == "__main__" guard around top-level work, that re-import would +# restart training in every worker. +if __name__ == "__main__": + # Data directory specification + tutorials_dir = Path(__file__).resolve().parent + # Fitted SSM surfaces and PCA coefficients written by Tutorial 8 (Duke Heart). + data_dir = tutorials_dir / "output" / "tutorial_08_duke_heart" + # PCA mean surface written by Tutorial 6 (Duke Heart); pca_model.json must + # sit beside it, which is how Tutorial 6 writes them. + ssm_mean_surface_file = DUKE_HEART.pca_mean_file + # Manifests, targets and the held-out evaluation are written here. + output_dir = tutorials_dir / "output" / "tutorial_09_duke_heart_mgn" + manifests_dir = output_dir / "manifests_mgn" + eval_dir = output_dir / "eval_mgn" + # The trained network is kept with the other trained networks instead. + model_output_dir = DUKE_HEART.mgn_weights_dir + + # Warm-start from a previous run's checkpoint; None trains from scratch. + # When resuming, training writes to a fresh sibling directory (``..._1``). + resume_from: Optional[Path] = None + + # Training hyperparameters. The heart template carries + # ParametersDukeHeartLabelmaps.model_points vertices, an order of magnitude + # fewer than the lung template of Tutorial 9 (lung), so a larger mini-batch + # fits; lower batch_size, or call + # training_method.set_num_processor_checkpoint_segments(...), on a smaller + # card. + epochs = 1500 + batch_size = 8 # mini-batch measured in (case, frame) graphs + learning_rate = 1.0e-3 + processor_size = 3 # message-passing hops + hidden_dim = 128 + num_layers = 2 # MLP layers inside each encoder / processor / decoder block + + # Explicit held-out splits; every other discovered case is used for + # training. The held-out case is the one Tutorials 2, 6 and 7 also keep + # out, so one patient stays unseen by everything in the pipeline. Adding a + # case to val_cases spends it on the intermittent validation RMSE instead of + # training; empty means that RMSE is reported as "n/a". + test_cases = [DUKE_HEART.hold_out_case] + val_cases: list[str] = [] + log_level = logging.INFO + + class_name = "tutorial_09_duke_heart_train_physicsnemo_mgn" + logging.basicConfig(level=log_level) + logger = logging.getLogger(class_name) + + # In test mode, train for a couple of epochs to keep the run tractable. + test_mode = TestTools.running_as_test() + if test_mode: + epochs = 2 + + if not ssm_mean_surface_file.exists(): + raise FileNotFoundError( + f"Tutorial 6 PCA mean surface not found: {ssm_mean_surface_file}\n" + "Run tutorials/tutorial_06_duke_heart_create_statistical_model.py first." + ) + + # Step 1: build one manifest per valid case and partition into splits. + manifests: dict[str, Path] = {} + for case_dir in sorted( + p for p in data_dir.glob("pm[0-9][0-9][0-9][0-9]") if p.is_dir() + ): + manifest_path = _write_case_manifest(case_dir, manifests_dir, logger) + if manifest_path is not None: + manifests[case_dir.name] = manifest_path + + if len(manifests) < 3: + raise RuntimeError( + f"Found only {len(manifests)} valid case(s) under {data_dir}; need at " + "least 3 to hold one out and still train on a population. See the " + "skip reasons logged above, and run " + "tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py first." + ) + + unknown = [ + case_id for case_id in test_cases + val_cases if case_id not in manifests + ] + if unknown: + raise ValueError(f"Split cases not found: {unknown}") + + test_manifests = [manifests[case_id] for case_id in test_cases] + val_manifests = [manifests[case_id] for case_id in val_cases] + train_manifests = [ + manifest_path + for case_id, manifest_path in manifests.items() + if case_id not in test_cases and case_id not in val_cases + ] + logger.info( + "Case split - train: %d, val: %d, test: %d", + len(train_manifests), + len(val_manifests), + len(test_manifests), + ) + + # Step 2: train the MeshGraphNet. The training method carries the network and + # its hyper-parameters; the workflow feeds it manifests and saves the results. + training_method = TrainPhysicsNeMoMGN(log_level=log_level) + training_method.set_epochs(epochs) + training_method.set_batch_size(batch_size) + training_method.set_learning_rate(learning_rate) + training_method.set_processor_size(processor_size) + training_method.set_hidden_dim(hidden_dim) + training_method.set_num_layers(num_layers) + + train_workflow = WorkflowTrainPhysicsNeMo( + train_manifests=train_manifests, + val_manifests=val_manifests, + pca_mean_mesh=ssm_mean_surface_file, + output_directory=model_output_dir, + resume_from=resume_from, + training_method=training_method, + log_level=log_level, + ) + train_result = train_workflow.process() + + # Step 3: evaluate held-out test cases against their ground-truth frames. + # When resuming, training writes to a fresh sibling directory, so evaluate + # the model from the directory training actually used. + model_directory = train_result["output_directory"] + infer_workflow = WorkflowInferPhysicsNeMo( + model_directory=model_directory, log_level=log_level + ) + # The targets are displacements from each case's reference surface, so the + # raw predictions are turned back into surfaces by the displacement decoder. + displacement_workflow = WorkflowInferMovement(infer_workflow, log_level=log_level) + + tutorial_results: dict[str, Any] = { + "model_directory": model_directory, + "cases": {}, + } + for case_id in test_cases: + logger.info("Evaluating held-out case %s", case_id) + tutorial_results["cases"][case_id] = displacement_workflow.process( + manifests[case_id], + output_directory=eval_dir / case_id, + ) + + # Testing: render the first predicted surface of the last held-out case and + # the RMSE-colored reference surface beside it. + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + baselines_dir=tutorials_dir.parent / "tests" / "baselines" / class_name, + log_level=log_level, + ) + last_case = tutorial_results["cases"][test_cases[-1]] + tutorial_results["screenshots"] = [ + tt.save_screenshot_mesh( + cast(pv.DataSet, pv.read(str(last_case["predicted_surfaces"][0]))), + "predicted_surface.png", + camera_position="iso", + color="limegreen", + ), + tt.save_screenshot_mesh( + cast(pv.DataSet, pv.read(str(last_case["rmse_surface"]))), + "rmse_surface.png", + camera_position="iso", + color="orange", + ), + ] diff --git a/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py b/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py index 24a14252..3a713833 100644 --- a/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py +++ b/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py @@ -58,19 +58,22 @@ Outputs ------- -Manifests are written under ``output/tutorial_09_lung_mgn/manifests_mgn/``: +Manifests, the held-out evaluation and the screenshots are written under +``output/tutorial_09_lung_mgn/``: - * ``Case*Pack_manifest.json`` - per-case training manifest - * ``Case*Pack_T??_ssm_surface_target.vtp`` - per-phase displacement targets + * ``manifests_mgn/Case*Pack_manifest.json`` - per-case training manifest + * ``manifests_mgn/Case*Pack_T??_ssm_surface_target.vtp`` - displacement targets + * ``eval_mgn/Case*Pack/`` - predicted surfaces per held-out case + * ``predicted_surface.png`` / ``rmse_surface.png`` - screenshots -The model and its evaluation land in the directory training actually used — -``output/tutorial_09_lung_mgn/`` normally, or a fresh ``..._2`` sibling when -resuming (see ``resume_from``), which is what ``tutorial_results`` reports as -``model_directory``: +The model lands in ``ParametersLungCTDirLab.mgn_weights_dir`` +(``network_weights/physicsnemo_mgn_lung_motion/``), where Tutorial 10 reads it: * ``mgn_stage_model.pt`` - trained MeshGraphNet checkpoint - * ``eval_mgn/Case*Pack/`` - predicted surfaces per held-out case - * ``predicted_surface.png`` / ``rmse_surface.png`` - screenshots + +Resuming (see ``resume_from``) writes the model to a fresh ``..._1`` sibling of +that directory instead, which is what ``tutorial_results`` reports as +``model_directory`` and what Tutorial 10 then has to be pointed at. """ # Imports @@ -83,6 +86,7 @@ import numpy as np import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( TestTools, @@ -183,13 +187,16 @@ def _write_case_manifest( ssm_mean_surface_file = ( tutorials_dir / "output" / "tutorial_06_lung" / "pca_mean_surface.vtp" ) - # All outputs (manifests, checkpoints, evaluation surfaces) are written here. + # Manifests, evaluation surfaces and screenshots are written here. output_dir = tutorials_dir / "output" / "tutorial_09_lung_mgn" manifests_dir = output_dir / "manifests_mgn" + # The trained model goes to the shared weights directory Tutorial 10 loads + # it from, beside the ICON weights the registration tutorials finetune. + weights_dir = LUNG_CT_DIRLAB.mgn_weights_dir # Warm-start from a previous run's checkpoint; None trains from scratch. When - # resuming, training writes to a fresh sibling directory (``..._2``), e.g. - # tutorials/output/tutorial_09_lung_mgn_2/mgn_stage_model_epoch_00200.pt + # resuming, training writes to a fresh sibling of weights_dir, e.g. + # network_weights/physicsnemo_mgn_lung_motion_1/mgn_stage_model_epoch_00200.pt resume_from: Optional[Path] = None # Training hyperparameters @@ -201,10 +208,11 @@ def _write_case_manifest( num_layers = 2 # MLP layers inside each encoder / processor / decoder block # Explicit held-out splits; every other discovered case is used for training. - # Case1 is also the case held out of the Tutorial 2 ICON finetuning. Adding a - # case to val_cases spends it on the intermittent validation RMSE instead of - # training; empty means that RMSE is reported as "n/a". - test_cases = ["Case1Pack"] + # The held-out case is the one Tutorial 10 predicts, and is also the case held + # out of the Tutorial 2 ICON finetuning. Adding a case to val_cases spends it + # on the intermittent validation RMSE instead of training; empty means that + # RMSE is reported as "n/a". + test_cases = [LUNG_CT_DIRLAB.mgn_hold_out_case] val_cases: list[str] = [] log_level = logging.INFO @@ -274,7 +282,7 @@ def _write_case_manifest( train_manifests=train_manifests, val_manifests=val_manifests, pca_mean_mesh=ssm_mean_surface_file, - output_directory=output_dir, + output_directory=weights_dir, resume_from=resume_from, training_method=training_method, log_level=log_level, @@ -283,7 +291,7 @@ def _write_case_manifest( # Step 3: evaluate held-out test cases against their ground-truth phases. # When resuming, training writes to a fresh sibling directory, so evaluate the - # model from the directory training actually used, not the original output_dir. + # model from the directory training actually used, not the original weights_dir. model_directory = train_result["output_directory"] infer_workflow = WorkflowInferPhysicsNeMo( model_directory=model_directory, log_level=log_level @@ -300,14 +308,14 @@ def _write_case_manifest( logger.info("Evaluating held-out case %s", case_id) tutorial_results["cases"][case_id] = displacement_workflow.process( manifests[case_id], - output_directory=model_directory / "eval_mgn" / case_id, + output_directory=output_dir / "eval_mgn" / case_id, ) # Testing: render the first predicted surface of the last held-out case and # the RMSE-colored reference surface beside it. tt = TestTools( class_name=class_name, - results_dir=model_directory, + results_dir=output_dir, baselines_dir=tutorials_dir.parent / "tests" / "baselines" / class_name, log_level=log_level, ) diff --git a/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py b/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py new file mode 100644 index 00000000..acb6ba93 --- /dev/null +++ b/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py @@ -0,0 +1,202 @@ +""" +Tutorial 10 (Duke Heart, MGN): Predict a Heart Surface at One Cardiac Stage + +Purpose +------- +Final stage of the Duke heart 4D deep-learning pipeline (Tutorials 8 -> 9 -> 10), +the counterpart of ``tutorial_10_lung_infer_physicsnemo_mgn.py``. A thin driver +over :class:`physiotwin4d.WorkflowInferPhysicsNeMo` and its displacement decoder +:class:`physiotwin4d.WorkflowInferMovement`: + +1. Discover the per-frame SSM surfaces that Tutorial 8 (Duke Heart) + (``tutorial_08_duke_heart_fit_model_to_4d_patients.py``) wrote for the test + case, and pick the cardiac stage to predict. Stages are parsed from the + ``g{PPP}`` gate tag of the frame filenames. + +2. Predict that case's surface at the chosen stage with the MeshGraphNet trained + by Tutorial 9 (``tutorial_09_duke_heart_train_physicsnemo_mgn.py``). The + network predicts per-vertex displacements, so the decoder adds them to the + case's reference SSM surface and scores the result in millimetres against the + ground-truth frame surface. + +3. Write the predicted surface as a USD (``WorkflowConvertVTKToUSD``, colored + with the heart anatomy material). + +For command-line use with path arguments, use the installed +``physiotwin4d-infer-physicsnemo`` CLI instead of editing this script. + +Extra Install Required +---------------------- +PhysicsNeMo and PyTorch Geometric must be installed:: + + pip install "physiotwin4d[physicsnemo]" + +Data Required +------------- + * ``output/tutorial_08_duke_heart//`` - Tutorial 8 SSM surfaces + * ``network_weights/physicsnemo_mgn_duke_heart_motion/mgn_stage_model.pt`` + - Tutorial 9 checkpoint + (``ParametersDukeHeartLabelmaps.mgn_weights_dir``) + +Outputs (under ``output/tutorial_10_duke_heart_mgn//``) +------------------------------------------------------------ + * ``_ssm_pca_coefficients_pred_s{TTT}.vtp`` - predicted surface + * ``_mgn_s{TTT}.usd`` - USD of that surface +""" + +# Imports +from __future__ import annotations + +import logging +from pathlib import Path +from typing import Any, Optional, cast + +import pyvista as pv +from parameters_duke_heart_labelmaps import DUKE_HEART + +from physiotwin4d import ( + TestTools, + WorkflowConvertVTKToUSD, + WorkflowInferMovement, + WorkflowInferPhysicsNeMo, +) + +# Gated frames carry a ``g{PPP}`` tag naming their percentage of the R-R +# interval; this is what a per-frame SSM surface is matched and staged by. +PHASE_SURFACE_PATTERN = "*_g[0-9][0-9][0-9]_*_ssm_surface.vtp" + + +def _cardiac_stage_from_filename(surface_file: Path) -> float: + """Extract the normalized cardiac stage [0, 1] from a ``g{PPP}`` filename stem.""" + for part in surface_file.stem.split("_"): + if part.startswith("g") and part[1:].isdigit(): + return int(part[1:]) / 100.0 + raise ValueError(f"Cannot parse cardiac gate from filename: {surface_file}") + + +# Only run if this script is not imported as a module + +# PhysicsNeMo and torch spawn worker processes. On Windows the spawn start +# method re-imports this script in each child; without the +# __name__ == "__main__" guard around top-level work, that re-import would +# restart the prediction in every worker. +if __name__ == "__main__": + # Data directory specification + tutorials_dir = Path(__file__).resolve().parent + # Fitted SSM surfaces and PCA coefficients written by Tutorial 8 (Duke Heart). + data_dir = tutorials_dir / "output" / "tutorial_08_duke_heart" + # The network Tutorial 9 (Duke Heart) trained. + model_dir = DUKE_HEART.mgn_weights_dir + # Intermittent-checkpoint epoch to load; None uses the final weights. + epoch: Optional[int] = None + + # Case to predict; the held-out test case of Tutorial 9 (Duke Heart). + case_id = DUKE_HEART.hold_out_case + # Fraction through the case's ordered gated frames to predict. + stage_fraction = 0.7 + + output_dir = tutorials_dir / "output" / "tutorial_10_duke_heart_mgn" / case_id + log_level = logging.INFO + + class_name = "tutorial_10_duke_heart_infer_physicsnemo" + logging.basicConfig(level=log_level) + logger = logging.getLogger(class_name) + + # Directory setup and data reading + + output_dir.mkdir(parents=True, exist_ok=True) + + checkpoint_file = model_dir / "mgn_stage_model.pt" + if not checkpoint_file.exists(): + raise FileNotFoundError( + f"Tutorial 9 checkpoint not found: {checkpoint_file}\n" + "Run tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py first." + ) + + case_dir = data_dir / case_id + reference_file = case_dir / f"{case_id}_ssm_surface.vtp" + pca_file = case_dir / f"{case_id}_ssm_pca_coefficients.json" + phase_files = sorted(case_dir.glob(PHASE_SURFACE_PATTERN)) + for required_file in (reference_file, pca_file): + if not required_file.exists(): + raise FileNotFoundError( + f"Tutorial 8 output not found: {required_file}\n" + "Run tutorials/tutorial_08_duke_heart_fit_model_to_4d_patients.py " + "first." + ) + if not phase_files: + raise FileNotFoundError(f"No gated frame surfaces found in {case_dir}") + + # Step 1: pick the test frame - the one 70% of the way through the case's + # ordered gated frames - and read its stage and ground-truth surface. + stages = [_cardiac_stage_from_filename(f) for f in phase_files] + test_index = int(stage_fraction * len(stages)) + test_stage = stages[test_index] + ground_truth_file = phase_files[test_index] + logger.info( + "Case %s: predicting stage %.2f (%s) of %d frames", + case_id, + test_stage, + ground_truth_file.name, + len(stages), + ) + + # Step 2: predict the case's surface at that stage with the trained + # MeshGraphNet. The model predicts displacements, so the displacement + # decoder adds them to the case's reference SSM surface and scores the + # result against the ground-truth frame surface in millimetres. + infer_workflow = WorkflowInferPhysicsNeMo( + model_directory=model_dir, epoch=epoch, log_level=log_level + ) + infer_result = WorkflowInferMovement( + infer_workflow, log_level=log_level + ).predict_single( + shape_parameters=pca_file, + stage=test_stage, + reference_mesh=reference_file, + ground_truth=ground_truth_file, + output_directory=output_dir, + ) + + # Step 3: write the predicted surface as a USD, colored with the heart + # anatomy material via USDAnatomyTools (appearance="anatomy"). The SSM is + # one structure, the whole heart minus its chamber cavities, so the surface + # is kept whole rather than split by connectivity. + usd_workflow = WorkflowConvertVTKToUSD( + input_meshes=[pv.read(str(infer_result["predicted_surface"]))], + usd_project_name=f"{case_id}_mgn_s{int(test_stage * 100):03d}", + output_directory=output_dir, + appearance="anatomy", + anatomy_type="heart", + separate_by_connectivity=False, + log_level=log_level, + ) + usd_file = usd_workflow.process()["usd_file"] + + tutorial_results: dict[str, Any] = dict(infer_result) + tutorial_results["stage"] = test_stage + tutorial_results["ground_truth_file"] = ground_truth_file + tutorial_results["usd_file"] = usd_file + + # Testing: render the predicted surface beside the ground-truth frame it is + # scored against. + tt = TestTools( + class_name=class_name, + results_dir=output_dir, + baselines_dir=tutorials_dir.parent / "tests" / "baselines" / class_name, + log_level=log_level, + ) + tutorial_results["screenshots"] = [ + tt.save_screenshot_mesh( + cast(pv.DataSet, pv.read(str(infer_result["predicted_surface"]))), + "predicted_surface.png", + camera_position="iso", + color="limegreen", + ), + tt.save_screenshot_mesh( + cast(pv.DataSet, pv.read(str(ground_truth_file))), + "ground_truth_surface.png", + camera_position="iso", + color="steelblue", + ), + ] diff --git a/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py b/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py index 08c3da03..3d7e5148 100644 --- a/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py +++ b/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py @@ -8,9 +8,11 @@ displacement decoder :class:`physiotwin4d.WorkflowInferMovement`: 1. Discover the per-phase SSM surfaces that Tutorial 8 - (``tutorial_08_lung_fit_model_to_4d_patients.py``) wrote for the test case, - and pick the respiratory stage to predict. Stages are parsed from the - ``T{PP}`` phase filenames. + (``tutorial_08_lung_fit_model_to_4d_patients.py``) wrote for + ``ParametersLungCTDirLab.mgn_hold_out_case`` -- the case Tutorial 9 held out + of training, so this scores generalization rather than recall -- and pick the + respiratory stage to predict. Stages are parsed from the ``T{PP}`` phase + filenames. 2. Predict that case's surface at the chosen stage with the MeshGraphNet trained by Tutorial 9 (``tutorial_09_lung_train_physicsnemo_mgn.py``). The @@ -32,13 +34,14 @@ Data Required ------------- - * ``output/tutorial_08_lung/Case1Pack/`` - Tutorial 8 SSM surfaces - * ``output/tutorial_09_lung_mgn/mgn_stage_model.pt`` - Tutorial 9 checkpoint - -Outputs (under ``output/tutorial_09_lung_mgn/tutorial_10_lung_mgn/Case1Pack/``) ------------------------------------------------------------------------------- - * ``Case1Pack_ssm_pca_coefficients_pred_s{TTT}.vtp`` - predicted surface - * ``Case1Pack_mgn_s{TTT}.usd`` - USD of that surface + * ``output/tutorial_08_lung//`` - Tutorial 8 SSM surfaces + * ``network_weights/physicsnemo_mgn_lung_motion/mgn_stage_model.pt`` + - Tutorial 9 checkpoint (``ParametersLungCTDirLab.mgn_weights_dir``) + +Outputs (under ``output/tutorial_10_lung_mgn//``) +------------------------------------------------------- + * ``_ssm_pca_coefficients_pred_s{TTT}.vtp`` - predicted surface + * ``_mgn_s{TTT}.usd`` - USD of that surface """ # Imports @@ -49,6 +52,7 @@ from typing import Any, Optional, cast import pyvista as pv +from parameters_lung_ct_dirlab import LUNG_CT_DIRLAB from physiotwin4d import ( TestTools, @@ -77,17 +81,18 @@ def _respiratory_stage_from_filename(surface_file: Path) -> float: tutorials_dir = Path(__file__).resolve().parent # Fitted SSM surfaces and PCA coefficients written by Tutorial 8 (lung). data_dir = tutorials_dir / "output" / "tutorial_08_lung" - # Tutorial 9 run directory to evaluate (matches that trainer's output_dir). - model_dir = tutorials_dir / "output" / "tutorial_09_lung_mgn" + # Weights Tutorial 9 trained. A resumed Tutorial 9 run writes to a numbered + # sibling of this directory, which is what would be evaluated instead. + model_dir = LUNG_CT_DIRLAB.mgn_weights_dir # Intermittent-checkpoint epoch to load; None uses the final weights. epoch: Optional[int] = None - # Case to predict; Case1Pack is the held-out test case of Tutorial 9. - case_id = "Case1Pack" + # Case to predict: the case Tutorial 9 held out of training. + case_id = LUNG_CT_DIRLAB.mgn_hold_out_case # Fraction through the case's ordered respiratory phases to predict. stage_fraction = 0.7 - output_dir = model_dir / "tutorial_10_lung_mgn" / case_id + output_dir = tutorials_dir / "output" / "tutorial_10_lung_mgn" / case_id log_level = logging.INFO class_name = "tutorial_10_lung_infer_physicsnemo_mgn" From a4a4b452d770a1d6f64955d443acba014639b8e5 Mon Sep 17 00:00:00 2001 From: Stephen Aylward Date: Tue, 11 Aug 2026 21:29:13 -0400 Subject: [PATCH 3/3] ENH: tutorial_04_duke_heart dir structure simplified. CodeRabbit edits. --- docs/tutorials.rst | 12 +-- pyproject.toml | 3 + src/physiotwin4d/contour_tools.py | 31 +++++-- src/physiotwin4d/train_physicsnemo_base.py | 11 ++- .../workflow_infer_physicsnemo.py | 7 +- .../workflow_train_physicsnemo.py | 83 ++++++++++++------- tests/test_contour_mesh_extraction.py | 10 +++ tests/test_image_tools.py | 4 +- tutorials/README.md | 10 ++- .../tutorial_04_duke_heart_labelmap_to_vtk.py | 16 +++- .../tutorial_05_duke_heart_vtk_to_usd.py | 31 ++++--- ..._06_duke_heart_create_statistical_model.py | 4 +- ..._heart_fit_statistical_model_to_patient.py | 28 ++++--- ...ial_09_duke_heart_train_physicsnemo_mgn.py | 11 ++- .../tutorial_09_lung_train_physicsnemo_mgn.py | 8 ++ ...utorial_10_duke_heart_infer_physicsnemo.py | 4 +- .../tutorial_10_lung_infer_physicsnemo_mgn.py | 7 +- 17 files changed, 196 insertions(+), 84 deletions(-) diff --git a/docs/tutorials.rst b/docs/tutorials.rst index 03e1d807..68e797b3 100644 --- a/docs/tutorials.rst +++ b/docs/tutorials.rst @@ -11,7 +11,9 @@ Tutorials

PhysioTwin4D tutorials

From a CT scan to an animated digital twin

- Ten numbered stages across 17 runnable Python scripts. + Ten numbered stages across 24 Python scripts, 16 of them runnable today: + the eight duke_heart variants wait on a dataset that is not + public yet. Each one drives the real workflow classes end-to-end on downloadable data, shows what it produced, and ends with the handful of constants to change so it runs on your own scans. @@ -230,10 +232,10 @@ Script it registers with a much tighter mask than the lungs, so its distance maps saturate over a shorter radius and do not share an intensity distribution with lung ones. The per-organ values live in - ``tutorials/parameters_lung_ct_dirlab.py`` and - ``tutorials/parameters_heart_ct_kcl.py``. This is a ``duke_heart`` tutorial: - Duke-Heart-4DLabelmaps is not publicly available yet, so it cannot be run — - see ``data/Duke-Heart-4DLabelmaps/README.md``. + ``tutorials/parameters_lung_ct_dirlab.py`` for the lung variant and + ``tutorials/parameters_duke_heart_labelmaps.py`` for this one. This is a + ``duke_heart`` tutorial: Duke-Heart-4DLabelmaps is not publicly available + yet, so it cannot be run — see ``data/Duke-Heart-4DLabelmaps/README.md``. Workflow :class:`~physiotwin4d.WorkflowFinetuneICONRegistration`, then diff --git a/pyproject.toml b/pyproject.toml index b75d518c..a4d94979 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -349,8 +349,11 @@ module = [ "tutorial_07_duke_heart_fit_statistical_model_to_patient", "tutorial_07_heart_fit_statistical_model_to_patient", "tutorial_07_lung_fit_statistical_model_to_patient", + "tutorial_08_duke_heart_fit_model_to_4d_patients", "tutorial_08_lung_fit_model_to_4d_patients", + "tutorial_09_duke_heart_train_physicsnemo_mgn", "tutorial_09_lung_train_physicsnemo_mgn", + "tutorial_10_duke_heart_infer_physicsnemo", "tutorial_10_lung_infer_physicsnemo_mgn", ] disable_error_code = ["import-not-found", "import-untyped"] diff --git a/src/physiotwin4d/contour_tools.py b/src/physiotwin4d/contour_tools.py index c0539301..dc35c4ac 100644 --- a/src/physiotwin4d/contour_tools.py +++ b/src/physiotwin4d/contour_tools.py @@ -298,8 +298,10 @@ def extract_label_surfaces( smoothing_iterations: Taubin smoothing iterations (0 disables). Returns: - Label id → that label's closed, outward-oriented surface. Empty if - the labelmap holds no non-zero label. + Label id → that label's closed, outward-oriented surface. A label + too small to survive the isotropic grid is left out, so the mapping + is empty when the labelmap holds no non-zero label and may be + missing labels that it does hold. """ labels = itk.GetArrayViewFromImage(labelmap_image) label_ids = [int(value) for value in np.unique(labels) if value != 0] @@ -369,6 +371,16 @@ def extract_label_surfaces( for label_id in label_ids: cell_ids: list[int] = np.flatnonzero(merged_ids == label_id).tolist() surface = self.extract_surface(merged.extract_cells(cell_ids)).triangulate() + if surface.n_cells == 0: + # A label smaller than the isotropic grid loses its vote to its + # neighbors, so no voxel is assigned to it and its field never + # crosses zero. It has no surface to return. + self.log_warning( + "Label %d is too small for a %.3g mm grid; it has no surface", + label_id, + iso, + ) + continue # The bookkeeping arrays of the merge and the split; the label is # the key of the returned mapping, so it is not data on the mesh. for array_name in ("LabelId", "vtkOriginalCellIds", "vtkOriginalPointIds"): @@ -490,8 +502,14 @@ def _signed_distance_mm( @staticmethod def is_watertight(surface: pv.PolyData) -> bool: - """Report whether every edge of *surface* is shared by exactly two faces.""" + """Report whether every edge of *surface* is shared by exactly two faces. + + A surface with no faces has no edge that fails the test, so it is + reported as not watertight rather than vacuously watertight. + """ faces = surface.triangulate().faces.reshape(-1, 4)[:, 1:] + if len(faces) == 0: + return False edges = np.sort( np.vstack([faces[:, [0, 1]], faces[:, [1, 2]], faces[:, [2, 0]]]), axis=1 ) @@ -524,10 +542,13 @@ def extract_tetrahedra( specific first. ``None`` leaves the mesh uncolored. Returns: - The structure's tetrahedral mesh, empty if *element_size_mm* is too - coarse to keep any of the mask. + The structure's tetrahedral mesh, empty if the mask is empty or + *element_size_mm* is too coarse to keep any of it. """ mask_arr = itk.GetArrayViewFromImage(mask_image) != 0 + if not mask_arr.any(): + self.log_warning("Mask holds no voxel to mesh; its mesh is empty") + return pv.UnstructuredGrid() # mask_arr axes are reversed relative to the ITK image, so the per-axis # extents come back as (z, y, x) and are flipped to (x, y, z). starts, stops = [], [] diff --git a/src/physiotwin4d/train_physicsnemo_base.py b/src/physiotwin4d/train_physicsnemo_base.py index 8fc38f42..db67fde1 100644 --- a/src/physiotwin4d/train_physicsnemo_base.py +++ b/src/physiotwin4d/train_physicsnemo_base.py @@ -116,7 +116,12 @@ def checkpoint_fields(self) -> dict: raise NotImplementedError def save_artifacts(self, output_dir: Path) -> None: - """Save any architecture-specific artifacts (MGN graph tensors).""" + """Save any architecture-specific artifacts (MGN graph tensors). + + Called by :meth:`train` as soon as :meth:`setup_inputs` has run, so the + artifacts are in place for inference from the first intermittent + checkpoint rather than only after the last epoch. + """ raise NotImplementedError # ─────────────────────────── Training loop ───────────────────────────── @@ -165,6 +170,10 @@ def train( self.log_info("Loaded model weights from %s", resume_from) self.setup_inputs(device, template_mesh, template_coords) + # Written now rather than after the last epoch: inference against an + # intermittent checkpoint needs them, and that is the point of writing + # those checkpoints while a long run is still going. + self.save_artifacts(output_dir) if sys.platform != "win32": try: diff --git a/src/physiotwin4d/workflow_infer_physicsnemo.py b/src/physiotwin4d/workflow_infer_physicsnemo.py index 488e6d9f..c153f5c5 100644 --- a/src/physiotwin4d/workflow_infer_physicsnemo.py +++ b/src/physiotwin4d/workflow_infer_physicsnemo.py @@ -87,7 +87,12 @@ def __init__( self.model_directory = Path(model_directory) tag = self.inference_method.model_tag - checkpoint_file = self.model_directory / f"{tag}_stage_model.pt" + if epoch is not None: + checkpoint_file = ( + self.model_directory / f"{tag}_stage_model_epoch_{epoch:05d}.pt" + ) + else: + checkpoint_file = self.model_directory / f"{tag}_stage_model.pt" if not checkpoint_file.exists(): raise FileNotFoundError(f"Model checkpoint not found: {checkpoint_file}") diff --git a/src/physiotwin4d/workflow_train_physicsnemo.py b/src/physiotwin4d/workflow_train_physicsnemo.py index 6faa73c9..937c12ae 100644 --- a/src/physiotwin4d/workflow_train_physicsnemo.py +++ b/src/physiotwin4d/workflow_train_physicsnemo.py @@ -189,6 +189,11 @@ def process(self) -> dict[str, Any]: stats["target_scale"], ) + # Everything inference needs except the weights, written before the + # first epoch so a run in progress can be evaluated from one of its + # intermittent checkpoints. + self._save_shared_assets(subjects, stats, output_dir, epochs) + model, losses, rmse_log = self.training_method.train( train_dataset, val_dataset, @@ -200,7 +205,7 @@ def process(self) -> dict[str, Any]: self._template_coords, resume_from=self.resume_from, ) - self._save_model(model, subjects, stats, losses, rmse_log, output_dir, epochs) + self._save_model(model, subjects, stats, losses, rmse_log, output_dir) self.log_section("PHYSICSNEMO %s TRAINING COMPLETE", model_tag.upper()) return { @@ -400,34 +405,26 @@ def _samples(split: str) -> list[_Sample]: ) return train_dataset, val_dataset - def _save_model( + def _save_shared_assets( self, - model: Any, subjects: dict[str, dict], stats: dict, - losses: list[float], - rmse_log: list[dict], output_dir: Path, epochs: int, ) -> None: - """Persist the checkpoint, metadata, logs and shared PCA assets.""" - import torch - + """Write the metadata and PCA assets inference needs beside the weights. + + None of this depends on the trained weights, so it is written before + training starts: + :class:`physiotwin4d.WorkflowInferPhysicsNeMo` reads the template mesh + and — through the inference method — the shared graph tensors from the + model directory, and cannot load an intermittent epoch checkpoint until + they are there. The training method writes its own artifacts once its + inputs are set up, at the top of its training loop. + """ method = self.training_method in_features = 3 + int(stats["pca_mean"].shape[0]) + 1 - tag = method.model_tag - checkpoint_file = output_dir / f"{tag}_stage_model.pt" - metadata_file = output_dir / f"{tag}_stage_model_metadata.json" - - train_ids = sorted(s for s, d in subjects.items() if d["split"] == "train") - val_ids = sorted(s for s, d in subjects.items() if d["split"] == "val") - - checkpoint = method.build_checkpoint(model, stats) - checkpoint["target_array"] = next(iter(subjects.values()))["target_array"] - checkpoint["train_subject_ids"] = train_ids - checkpoint["val_subject_ids"] = val_ids - checkpoint["resumed_from"] = str(self.resume_from) if self.resume_from else None - torch.save(checkpoint, checkpoint_file) + metadata_file = output_dir / f"{method.model_tag}_stage_model_metadata.json" n_pca = int(stats["pca_mean"].shape[0]) n_target = int(stats["n_target"]) @@ -458,6 +455,40 @@ def _save_model( metadata.update(method.checkpoint_fields()) metadata_file.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + # Copy PCA assets so the model directory is self-contained for inference. + shutil.copy2(self.pca_mean_mesh, output_dir / self.pca_mean_mesh.name) + suffix = ".vtp" if isinstance(self._template_mesh, pv.PolyData) else ".vtu" + self._template_mesh.save(str(output_dir / f"pca_mean_template{suffix}")) + if self._pca_model_path is not None: + shutil.copy2(self._pca_model_path, output_dir / "pca_model.json") + + self.metadata_file = metadata_file + + def _save_model( + self, + model: Any, + subjects: dict[str, dict], + stats: dict, + losses: list[float], + rmse_log: list[dict], + output_dir: Path, + ) -> None: + """Persist the final checkpoint and the training logs.""" + import torch + + tag = self.training_method.model_tag + checkpoint_file = output_dir / f"{tag}_stage_model.pt" + + train_ids = sorted(s for s, d in subjects.items() if d["split"] == "train") + val_ids = sorted(s for s, d in subjects.items() if d["split"] == "val") + + checkpoint = self.training_method.build_checkpoint(model, stats) + checkpoint["target_array"] = next(iter(subjects.values()))["target_array"] + checkpoint["train_subject_ids"] = train_ids + checkpoint["val_subject_ids"] = val_ids + checkpoint["resumed_from"] = str(self.resume_from) if self.resume_from else None + torch.save(checkpoint, checkpoint_file) + (output_dir / "training_losses.json").write_text( json.dumps(losses, indent=2), encoding="utf-8" ) @@ -471,17 +502,7 @@ def _save_model( writer.writeheader() writer.writerows(rmse_log) - # Copy PCA assets so the model directory is self-contained for inference. - shutil.copy2(self.pca_mean_mesh, output_dir / self.pca_mean_mesh.name) - suffix = ".vtp" if isinstance(self._template_mesh, pv.PolyData) else ".vtu" - self._template_mesh.save(str(output_dir / f"pca_mean_template{suffix}")) - if self._pca_model_path is not None: - shutil.copy2(self._pca_model_path, output_dir / "pca_model.json") - - method.save_artifacts(output_dir) - self.checkpoint_file = checkpoint_file - self.metadata_file = metadata_file self.training_loss = losses self.val_rmse_log = rmse_log self.log_info("Model saved to %s", checkpoint_file) diff --git a/tests/test_contour_mesh_extraction.py b/tests/test_contour_mesh_extraction.py index 862eb083..fb0e9891 100644 --- a/tests/test_contour_mesh_extraction.py +++ b/tests/test_contour_mesh_extraction.py @@ -263,6 +263,10 @@ def test_normals_point_outward(self, contour_tools: ContourTools) -> None: ) assert surface.volume > 0.0, "enclosed volume must be positive" + def test_empty_surface_is_not_watertight(self, contour_tools: ContourTools) -> None: + """A surface with no face fails the test rather than passing it vacuously.""" + assert not contour_tools.is_watertight(pv.PolyData()) + def test_decimation_reduces_triangle_count( self, contour_tools: ContourTools, heart_labelmap: itk.Image ) -> None: @@ -319,6 +323,12 @@ def test_volume_survives_coarsening( assert coarse_volume == pytest.approx(full_volume, rel=0.1) assert coarse.n_cells == pytest.approx(full.n_cells * 0.125, rel=0.3) + def test_empty_mask_yields_an_empty_mesh(self, contour_tools: ContourTools) -> None: + """A mask with nothing in it has no bounding box to mesh.""" + empty = itk.GetImageFromArray(np.zeros((4, 4, 4), dtype=np.uint8)) + + assert contour_tools.extract_tetrahedra(empty).n_cells == 0 + def test_elements_are_isotropic(self, contour_tools: ContourTools) -> None: """The requested element size is what the mesh is built on. diff --git a/tests/test_image_tools.py b/tests/test_image_tools.py index 006da590..d4614af9 100644 --- a/tests/test_image_tools.py +++ b/tests/test_image_tools.py @@ -526,7 +526,9 @@ def test_nearest_neighbor_keeps_input_values(self, image_tools: ImageTools) -> N out = image_tools.resample_image_by_scale(itk_image, 0.5, interpolate=False) - assert set(np.unique(itk.array_from_image(out))) <= {0.0, 7.0} + values = set(np.unique(itk.array_from_image(out))) + assert values <= {0.0, 7.0} + assert 7.0 in values, "coarsening must keep the block, not sample past it" @pytest.mark.parametrize("scale", [0.0, -1.0]) def test_non_positive_scale_raises( diff --git a/tutorials/README.md b/tutorials/README.md index 3245cf38..962dca89 100644 --- a/tutorials/README.md +++ b/tutorials/README.md @@ -100,8 +100,8 @@ variants for the anatomy you care about: every tutorial consumes the output of its own anatomy's earlier tutorials, never the other's. 1. **Tutorial 1** converts one gated 4D CT into an animated USD - the heart variant uses Slicer-Heart-CT, the lung variant DirLab-4DCT. Prepare the dataset for your anatomy per `data/README.md`, then start here. -2. **Tutorial 2** requires DirLab-4DCT (download it per `data/README.md`) and finetunes the ICON weights Tutorials 3 (lung) and 8 use when they are present — both fall back to the stock uniGradICON weights otherwise. -3. **Tutorial 3** registers with those weights; the heart variant uses Slicer-Heart-CT, the lung variant DirLab-4DCT. +2. **Tutorial 2** requires DirLab-4DCT (download it per `data/README.md`) and finetunes the ICON weights Tutorial 8 uses when they are present — it falls back to the stock uniGradICON weights otherwise. +3. **Tutorial 3** registers with Greedy and needs no finetuned weights; the heart variant uses Slicer-Heart-CT, the lung variant DirLab-4DCT. 4. **Tutorial 4** segments a CT into VTK surfaces; the heart variant uses Slicer-Heart-CT, the lung variant DirLab-4DCT. 5. **Tutorial 5** (heart only) uses the VTK surfaces produced by Tutorial 4 (heart) - run Tutorial 4 first. 6. **Tutorial 6** creates the PCA statistical model; the heart variant from KCL-Heart-Model, the lung variant from the DirLab-4DCT `Case*T70.mha` phases, which it segments itself. Both write `pca_model.json` and `pca_mean_surface.vtp` under their own output directory. @@ -114,6 +114,12 @@ Tutorial 6 lung model, in order: 9. **Tutorial 9** trains a PhysicsNeMo MeshGraphNet to predict the per-vertex motion at any stage. PhysicsNeMo is an optional extra: install with `pip install "physiotwin4d[physicsnemo]"` (requires Python >= 3.11); the MeshGraphNet also needs `torch-geometric`. A `TrainPhysicsNeMoMLP` method exists as a drop-in alternative, without its own tutorial. 10. **Tutorial 10** loads that checkpoint and predicts one case's surface at a requested stage, scoring it against the acquired phase and exporting USD. The case, checkpoint epoch, and stage are constants near the top of the script; for command-line runs with path arguments, use the installed `physiotwin4d-infer-physicsnemo` CLI. +The `duke_heart` variants form their own chain on Duke-Heart-4DLabelmaps, +which no step above shares: Tutorial 4 (duke heart) -> 5 -> 6 -> 7 -> 8 -> 9 -> +10, each reading the previous one's output, with Tutorial 2 (heart distancemap +variant) supplying optional finetuned weights to Tutorials 7 and 8. That +dataset is not publicly available yet, so this chain cannot be run today. + ## For Contributors Class-level API reference: [../docs/api/index.rst](../docs/api/index.rst) diff --git a/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py b/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py index e7340afb..ff50cf01 100644 --- a/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py +++ b/tutorials/tutorial_04_duke_heart_labelmap_to_vtk.py @@ -231,6 +231,9 @@ def write_tetrahedra( surface_count = 0 whole_heart_surface: Optional[pv.PolyData] = None + # Every case writes into this one directory: the frame stems already start + # with their case id, so the names stay unique and the readers downstream + # group on that prefix rather than on a directory. case_output_dir = output_dir case_output_dir.mkdir(parents=True, exist_ok=True) @@ -263,11 +266,20 @@ def write_tetrahedra( whole_heart_mask = mask_from( labelmap_image, np.isin(labels, whole_heart_ids) ) - whole_heart = contour_tools.extract_label_surfaces( + # The mask carries one label, so its surface is the only entry; an + # empty mapping means the frame held no heart to contour. + whole_heart_surfaces = contour_tools.extract_label_surfaces( whole_heart_mask, isotropic_spacing_mm=surface_spacing_mm, smoothing_iterations=smoothing_iterations, - )[1] + ) + if 1 not in whole_heart_surfaces: + reporter.log_warning( + "%s: no whole-heart surface; skipping the frame", + labelmap_file.name, + ) + continue + whole_heart = whole_heart_surfaces[1] contour_tools.apply_anatomy_color(whole_heart, [whole_heart_name, "heart"]) annotate(whole_heart, whole_heart_ids, whole_heart_name, labelmap_file.name) surface_count += 1 diff --git a/tutorials/tutorial_05_duke_heart_vtk_to_usd.py b/tutorials/tutorial_05_duke_heart_vtk_to_usd.py index 6c1eebc5..c7f0cc37 100644 --- a/tutorials/tutorial_05_duke_heart_vtk_to_usd.py +++ b/tutorials/tutorial_05_duke_heart_vtk_to_usd.py @@ -26,7 +26,7 @@ Data Required ------------- -``tutorials/output/tutorial_04_duke_heart_labelmap/pm????/*_surfaces.vtp`` +``tutorials/output/tutorial_04_duke_heart_labelmap/pm????_*_surfaces.vtp`` (run ``tutorial_04_duke_heart_labelmap_to_vtk.py`` first) """ @@ -66,13 +66,15 @@ output_dir.mkdir(parents=True, exist_ok=True) - case_dirs = [input_dir / "pm0004"] - # sorted( - # path for path in input_dir.glob("pm[0-9][0-9][0-9][0-9]") if path.is_dir() - # ) - if not case_dirs: + # Tutorial 4 writes every case's frames into one directory, and every frame + # stem starts with its case id, so the cases are the distinct prefixes. + # Sorting by name puts each case's frames in gating order. + frame_files: dict[str, list[Path]] = {} + for vtp_file in sorted(input_dir.glob("pm[0-9][0-9][0-9][0-9]_*_surfaces.vtp")): + frame_files.setdefault(vtp_file.name.split("_")[0], []).append(vtp_file) + if not frame_files: raise FileNotFoundError( - f"No pm???? case directories found under {input_dir}.\n" + f"No pm????_*_surfaces.vtp frame surfaces found under {input_dir}.\n" "Run tutorial_04_duke_heart_labelmap_to_vtk.py first." ) @@ -81,16 +83,11 @@ last_time_codes: list[float] = [] all_structures: set[str] = set() label_names = segmenter.taxonomy.all_labels() - for case_dir in case_dirs: - # One file per frame, already holding every structure; sorting by name - # puts the frames in gating order. + for case_id, case_files in frame_files.items(): + # One file per frame, already holding every structure. frame_meshes = [ - cast(pv.PolyData, pv.read(str(vtp_file))) - for vtp_file in sorted(case_dir.glob("*_surfaces.vtp")) + cast(pv.PolyData, pv.read(str(vtp_file))) for vtp_file in case_files ] - if not frame_meshes: - reporter.log_warning("%s: no per-frame surface files", case_dir.name) - continue # The merged files carry label ids but not names, which come from the # taxonomy of the segmenter that produced the labelmaps. @@ -101,13 +98,13 @@ } all_structures.update(mask_ids.values()) reporter.log_section( - f"{case_dir.name}: {len(frame_meshes)} frames, {len(mask_ids)} structures" + f"{case_id}: {len(frame_meshes)} frames, {len(mask_ids)} structures" ) # One frame per time code and one cardiac cycle per second of playback. workflow = WorkflowConvertVTKToUSD( input_meshes=frame_meshes, - usd_project_name=case_dir.name, + usd_project_name=case_id, output_directory=output_dir, separate_by_connectivity=False, appearance="anatomy", diff --git a/tutorials/tutorial_06_duke_heart_create_statistical_model.py b/tutorials/tutorial_06_duke_heart_create_statistical_model.py index 8e36b4a4..8ba277e3 100644 --- a/tutorials/tutorial_06_duke_heart_create_statistical_model.py +++ b/tutorials/tutorial_06_duke_heart_create_statistical_model.py @@ -17,7 +17,7 @@ Data Required ------------- -``tutorials/output/tutorial_04_duke_heart_labelmap/pm????/ +``tutorials/output/tutorial_04_duke_heart_labelmap/ *_ref_heart_minus_interior_chambers.vtp`` (run ``tutorial_04_duke_heart_labelmap_to_vtk.py`` first) @@ -116,7 +116,7 @@ if not reference_surface_file.exists(): mean_workflow = WorkflowCreateMeanSurface( surfaces=sample_surfaces, - template_surface=sample_surfaces[4], + template_surface=sample_surfaces[len(sample_surfaces) // 2], log_level=log_level, ) mean_workflow.set_number_of_iterations(mean_surface_iterations) diff --git a/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py b/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py index 37c89da2..98638761 100644 --- a/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py +++ b/tutorials/tutorial_07_duke_heart_fit_statistical_model_to_patient.py @@ -36,7 +36,7 @@ import json import logging from pathlib import Path -from typing import Any, Optional, cast +from typing import Any, cast import itk import numpy as np @@ -111,10 +111,15 @@ ) pca_mean = cast(pv.DataSet, pv.read(str(pca_mean_file))) - pca_model: Optional[dict[str, Any]] = None - if pca_json.exists(): - with pca_json.open(encoding="utf-8") as f: - pca_model = json.load(f) + # The fit is a PCA fit: without the model there is nothing to fit, and the + # workflow's PCA outputs read below would never be set. + if not pca_json.exists(): + raise FileNotFoundError( + f"Tutorial 6 PCA model not found: {pca_json}\n" + "Run tutorials/tutorial_06_duke_heart_create_statistical_model.py first." + ) + with pca_json.open(encoding="utf-8") as f: + pca_model: dict[str, Any] = json.load(f) patient_labelmap_files = sorted( (data_dir / patient_case).glob("*_ref_labelmap.nii.gz") @@ -162,13 +167,12 @@ ) workflow.set_mask_dilation_mm(DUKE_HEART.mask_dilation_mm) workflow.set_distancemap_squared_max(DUKE_HEART.distancemap_squared_max) - if pca_model is not None: - workflow.set_use_pca_registration( - use_pca_registration=True, - pca_model=pca_model, - number_of_pca_components=number_of_pca_components, - use_surface=False, - ) + workflow.set_use_pca_registration( + use_pca_registration=True, + pca_model=pca_model, + number_of_pca_components=number_of_pca_components, + use_surface=False, + ) # The labelmap-to-labelmap stage registers distance maps, not intensities, # so it uses the distance-map-finetuned weights when they exist; without diff --git a/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py b/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py index e319bc79..afae05f5 100644 --- a/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py +++ b/tutorials/tutorial_09_duke_heart_train_physicsnemo_mgn.py @@ -67,8 +67,15 @@ reports as ``model_directory``: * ``mgn_stage_model.pt`` - trained MeshGraphNet checkpoint - * ``pca_mean_surface.vtp``, ``pca_model.json`` - PCA assets copied beside it, - so the directory is self-contained for inference + * ``mgn_stage_model_epoch_#####.pt`` - intermittent checkpoints + * ``pca_mean_surface.vtp``, ``pca_mean_template.vtp``, ``pca_model.json``, + ``shared_edge_index.pt``, ``shared_edge_features.pt`` and the metadata JSON + - everything inference needs beside the weights + +Everything but the checkpoints is written before the first epoch, so +``tutorial_10_duke_heart_infer_physicsnemo.py`` can be pointed at this +directory with its ``epoch`` set to an intermittent checkpoint while training +is still running. """ # Imports diff --git a/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py b/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py index 3a713833..446d412a 100644 --- a/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py +++ b/tutorials/tutorial_09_lung_train_physicsnemo_mgn.py @@ -70,6 +70,14 @@ (``network_weights/physicsnemo_mgn_lung_motion/``), where Tutorial 10 reads it: * ``mgn_stage_model.pt`` - trained MeshGraphNet checkpoint + * ``mgn_stage_model_epoch_#####.pt`` - intermittent checkpoints + * ``pca_mean_surface.vtp``, ``pca_mean_template.vtp``, ``pca_model.json``, + ``shared_edge_index.pt``, ``shared_edge_features.pt`` and the metadata JSON + - everything inference needs beside the weights + +Everything but the checkpoints is written before the first epoch, so Tutorial 10 +can be run against an intermittent checkpoint (its ``epoch`` constant) while +this training run is still going. Resuming (see ``resume_from``) writes the model to a fresh ``..._1`` sibling of that directory instead, which is what ``tutorial_results`` reports as diff --git a/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py b/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py index acb6ba93..b6a3ad7e 100644 --- a/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py +++ b/tutorials/tutorial_10_duke_heart_infer_physicsnemo.py @@ -130,7 +130,9 @@ def _cardiac_stage_from_filename(surface_file: Path) -> float: # Step 1: pick the test frame - the one 70% of the way through the case's # ordered gated frames - and read its stage and ground-truth surface. stages = [_cardiac_stage_from_filename(f) for f in phase_files] - test_index = int(stage_fraction * len(stages)) + # Clamped, so a stage_fraction of 1.0 picks the last frame rather than one + # past it. + test_index = min(int(stage_fraction * len(stages)), len(stages) - 1) test_stage = stages[test_index] ground_truth_file = phase_files[test_index] logger.info( diff --git a/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py b/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py index 3d7e5148..f6acbdcd 100644 --- a/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py +++ b/tutorials/tutorial_10_lung_infer_physicsnemo_mgn.py @@ -85,7 +85,7 @@ def _respiratory_stage_from_filename(surface_file: Path) -> float: # sibling of this directory, which is what would be evaluated instead. model_dir = LUNG_CT_DIRLAB.mgn_weights_dir # Intermittent-checkpoint epoch to load; None uses the final weights. - epoch: Optional[int] = None + epoch: Optional[int] = 200 # Case to predict: the case Tutorial 9 held out of training. case_id = LUNG_CT_DIRLAB.mgn_hold_out_case @@ -103,7 +103,10 @@ def _respiratory_stage_from_filename(surface_file: Path) -> float: output_dir.mkdir(parents=True, exist_ok=True) - checkpoint_file = model_dir / "mgn_stage_model.pt" + if epoch is not None: + checkpoint_file = model_dir / f"mgn_stage_model_epoch_{epoch:05d}.pt" + else: + checkpoint_file = model_dir / "mgn_stage_model.pt" if not checkpoint_file.exists(): raise FileNotFoundError( f"Tutorial 9 checkpoint not found: {checkpoint_file}\n"