diff --git a/docs/changes/newsfragments/462.feature b/docs/changes/newsfragments/462.feature new file mode 100644 index 000000000..b709d4c83 --- /dev/null +++ b/docs/changes/newsfragments/462.feature @@ -0,0 +1 @@ +Enable :class:`.SpaceWarper` to warp data from native space to template spaces via both FSL and ANTs by `Synchon Mandal`_ diff --git a/junifer/preprocess/warping/_ants_warper.py b/junifer/preprocess/warping/_ants_warper.py index 3ecb5711c..7a94432b1 100644 --- a/junifer/preprocess/warping/_ants_warper.py +++ b/junifer/preprocess/warping/_ants_warper.py @@ -180,27 +180,74 @@ class ANTsWarper: # Template space warping else: + input_space = input["space"] logger.debug( - f"Using ANTs to warp data from {input['space']} to {reference}" + f"Using ANTs to warp data from {input_space} space " + f"to {reference} space" ) - # Get xfm file - xfm_file_path = get_xfm(src=input["space"], dst=reference) - # Get template space image - template_space_img = get_template( - space=reference, - target_img=input["data"], - extra_input=None, - ) - # Save template - template_space_img_path = ( - element_tempdir / f"{reference}_T1w.nii.gz" - ) - nib.save(template_space_img, template_space_img_path) + # Native to MNI + if input_space == "native": + # Get warp file path + xfm_file_path = None + for entry in extra_input["Warp"]: + if entry["src"] == "native" and entry["dst"] == reference: + xfm_file_path = entry["path"] + if xfm_file_path is None: + raise_error( + klass=RuntimeError, + msg="Could not find correct warp file path", + ) + + # Use ResampleImage if input data resolution and reference + # resolution don't match + input_res = np.min(input["data"].header.get_zooms()[:3]) + ref_res = np.min( + input["reference"]["data"].header.get_zooms()[:3] + ) + logger.debug(f"Input resolution: {input_res}") + logger.debug(f"Reference resolution: {ref_res}") + if input_res != ref_res: + # Create a tempfile for resampled reference output + ref_path = ( + element_tempdir + / f"resampled_reference-{reference}.nii.gz" + ) + # Set ResampleImage command + resample_image_cmd = [ + "ResampleImage", + "3", # image dimension + f"{input['reference']['path'].resolve()}", + f"{ref_path.resolve()}", + f"{input_res}x{input_res}x{input_res}", + "0", # option for spacing and not size + "3 3", # Lanczos windowed sinc + ] + # Call ResampleImage + run_ext_cmd(name="ResampleImage", cmd=resample_image_cmd) + else: + logger.debug( + "Reference resolution matches input resolution" + ) + ref_path = input["reference"]["path"] + + # MNI to MNI + else: + # Get xfm file + xfm_file_path = get_xfm(src=input_space, dst=reference) + # Get template space image in correct resolution + template_space_img = get_template( + space=reference, + target_img=input["data"], + extra_input=None, + ) + # Save template + ref_path = element_tempdir / f"{reference}_T1w.nii.gz" + nib.save(template_space_img, ref_path) # Create a tempfile for warped output warped_output_path = element_tempdir / ( - f"warped_data_from_{input['space']}_to_{reference}.nii.gz" + f"warped_data_from_{input_space}_to_{reference}.nii.gz" ) # Set antsApplyTransforms command @@ -210,7 +257,7 @@ class ANTsWarper: "-e 3", "-n LanczosWindowedSinc", f"-i {input['path'].resolve()}", - f"-r {template_space_img_path.resolve()}", + f"-r {ref_path.resolve()}", f"-t {xfm_file_path.resolve()}", f"-o {warped_output_path.resolve()}", ] @@ -226,18 +273,20 @@ class ANTsWarper: "data": nib.load(warped_output_path), # Update warped input's space "space": reference, - # Save reference path - "reference": {"path": template_space_img_path}, + # Save resampled reference path or overwrite original + # keeping it same + "reference": {"path": ref_path}, # Keep pre-warp space for further operations - "prewarp_space": input["space"], + "prewarp_space": input_space, } ) # Check for data type's mask and warp if found if input.get("mask") is not None: + logger.debug("Warping associated mask") # Create a tempfile for warped mask output apply_transforms_mask_out_path = element_tempdir / ( - f"warped_mask_from_{input['space']}_to_{reference}.nii.gz" + f"warped_mask_from_{input_space}_to_{reference}.nii.gz" ) # Set antsApplyTransforms command apply_transforms_mask_cmd = [ @@ -246,8 +295,8 @@ class ANTsWarper: "-e 3", "-n 'GenericLabel[NearestNeighbor]'", f"-i {input['mask']['path'].resolve()}", - # use resampled reference - f"-r {input['reference']['path'].resolve()}", + # use resampled reference or original + f"-r {ref_path.resolve()}", f"-t {xfm_file_path.resolve()}", f"-o {apply_transforms_mask_out_path.resolve()}", ] diff --git a/junifer/preprocess/warping/_fsl_warper.py b/junifer/preprocess/warping/_fsl_warper.py index 3edae21f2..51178b946 100644 --- a/junifer/preprocess/warping/_fsl_warper.py +++ b/junifer/preprocess/warping/_fsl_warper.py @@ -40,6 +40,7 @@ class FSLWarper: self, input: dict[str, Any], extra_input: dict[str, Any], + reference: str, ) -> dict[str, Any]: # pragma: no cover """Preprocess using FSL. @@ -50,6 +51,10 @@ class FSLWarper: extra_input : dict The other fields in the Junifer Data object. Should have ``T1w`` and ``Warp`` data types. + reference : str + The data type or template space to use as reference for warping. + Template space conversion is only possible from native space, + not from another template space. Returns ------- @@ -62,102 +67,223 @@ class FSLWarper: If warp file path could not be found in ``extra_input``. """ - logger.debug("Using FSL for space warping") - - # Get the min of the voxel sizes from input and use it as the - # resolution - resolution = np.min(input["data"].header.get_zooms()[:3]) - - # Get warp file path - warp_file_path = None - for entry in extra_input["Warp"]: - if entry["dst"] == "native": - warp_file_path = entry["path"] - if warp_file_path is None: - raise_error( - klass=RuntimeError, msg="Could not find correct warp file path" - ) - # Create element-specific tempdir for storing post-warping assets element_tempdir = WorkDirManager().get_element_tempdir( prefix="fsl_warper" ) - # Create a tempfile for resampled reference output - flirt_out_path = element_tempdir / "resampled_reference.nii.gz" - # Set flirt command - flirt_cmd = [ - "flirt", - "-interp spline", - f"-in {extra_input['T1w']['path'].resolve()}", - f"-ref {extra_input['T1w']['path'].resolve()}", - f"-applyisoxfm {resolution}", - f"-out {flirt_out_path.resolve()}", - ] - # Call flirt - run_ext_cmd(name="flirt", cmd=flirt_cmd) + # Warping to native space + if reference == "T1w": + logger.debug("Using FSL for space warping") - # Create a tempfile for warped output - applywarp_out_path = element_tempdir / "warped_data.nii.gz" - # Set applywarp command - applywarp_cmd = [ - "applywarp", - "--interp=spline", - f"-i {input['path'].resolve()}", - # use resampled reference - f"-r {flirt_out_path.resolve()}", - f"-w {warp_file_path.resolve()}", - f"-o {applywarp_out_path.resolve()}", - ] - # Call applywarp - run_ext_cmd(name="applywarp", cmd=applywarp_cmd) + # Get the min of the voxel sizes from input and use it as the + # resolution + resolution = np.min(input["data"].header.get_zooms()[:3]) - logger.debug("Updating warped data") - input.update( - { - # Update path to sync with "data" - "path": applywarp_out_path, - # Load nifti - "data": nib.load(applywarp_out_path), - # Use reference input's space as warped input's space - "space": extra_input["T1w"]["space"], - # Save resampled reference path - "reference": {"path": flirt_out_path}, - # Keep pre-warp space for further operations - "prewarp_space": input["space"], - } - ) + # Get warp file path + warp_file_path = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + warp_file_path = entry["path"] + if warp_file_path is None: + raise_error( + klass=RuntimeError, + msg="Could not find correct warp file path", + ) - # Check for data type's mask and warp if found - if input.get("mask") is not None: - # Create a tempfile for warped mask output - applywarp_mask_out_path = element_tempdir / "warped_mask.nii.gz" + # Create a tempfile for resampled reference output + flirt_out_path = element_tempdir / "resampled_reference.nii.gz" + # Set flirt command + flirt_cmd = [ + "flirt", + "-interp spline", + f"-in {extra_input['T1w']['path'].resolve()}", + f"-ref {extra_input['T1w']['path'].resolve()}", + f"-applyisoxfm {resolution}", + f"-out {flirt_out_path.resolve()}", + ] + # Call flirt + run_ext_cmd(name="flirt", cmd=flirt_cmd) + + # Create a tempfile for warped output + applywarp_out_path = element_tempdir / "warped_data.nii.gz" # Set applywarp command - applywarp_mask_cmd = [ + applywarp_cmd = [ "applywarp", - "--interp=nn", - f"-i {input['mask']['path'].resolve()}", + "--interp=spline", + f"-i {input['path'].resolve()}", # use resampled reference - f"-r {input['reference']['path'].resolve()}", + f"-r {flirt_out_path.resolve()}", f"-w {warp_file_path.resolve()}", - f"-o {applywarp_mask_out_path.resolve()}", + f"-o {applywarp_out_path.resolve()}", ] # Call applywarp - run_ext_cmd(name="applywarp", cmd=applywarp_mask_cmd) + run_ext_cmd(name="applywarp", cmd=applywarp_cmd) - logger.debug("Updating warped mask data") + logger.debug("Updating warped data") input.update( { - "mask": { - # Update path to sync with "data" - "path": applywarp_mask_out_path, - # Load nifti - "data": nib.load(applywarp_mask_out_path), - # Use reference input's space as warped input mask's - # space - "space": extra_input["T1w"]["space"], - } + # Update path to sync with "data" + "path": applywarp_out_path, + # Load nifti + "data": nib.load(applywarp_out_path), + # Use reference input's space as warped input's space + "space": extra_input["T1w"]["space"], + # Save resampled reference path + "reference": {"path": flirt_out_path}, + # Keep pre-warp space for further operations + "prewarp_space": input["space"], } ) + # Check for data type's mask and warp if found + if input.get("mask") is not None: + # Create a tempfile for warped mask output + applywarp_mask_out_path = ( + element_tempdir / "warped_mask.nii.gz" + ) + # Set applywarp command + applywarp_mask_cmd = [ + "applywarp", + "--interp=nn", + f"-i {input['mask']['path'].resolve()}", + # use resampled reference + f"-r {input['reference']['path'].resolve()}", + f"-w {warp_file_path.resolve()}", + f"-o {applywarp_mask_out_path.resolve()}", + ] + # Call applywarp + run_ext_cmd(name="applywarp", cmd=applywarp_mask_cmd) + + logger.debug("Updating warped mask data") + input.update( + { + "mask": { + # Update path to sync with "data" + "path": applywarp_mask_out_path, + # Load nifti + "data": nib.load(applywarp_mask_out_path), + # Use reference input's space as warped input + # mask's space + "space": extra_input["T1w"]["space"], + } + } + ) + + # Warping from native to template space + else: + logger.debug( + f"Using FSL to warp data from native space to {reference} " + "space" + ) + + # Get warp file path + warp_file_path = None + for entry in extra_input["Warp"]: + if entry["src"] == "native" and entry["dst"] == reference: + warp_file_path = entry["path"] + if warp_file_path is None: + raise_error( + klass=RuntimeError, + msg="Could not find correct warp file path", + ) + + # Use flirt if input data resolution and reference resolution don't + # match + input_resolution = np.min(input["data"].header.get_zooms()[:3]) + ref_resolution = np.min( + input["reference"]["data"].header.get_zooms()[:3] + ) + logger.debug(f"Input resolution: {input_resolution}") + logger.debug(f"Reference resolution: {ref_resolution}") + if input_resolution != ref_resolution: + logger.debug("Resampling reference to match input resolution") + # Create a tempfile for resampled reference output + ref_path = ( + element_tempdir / f"resampled_reference-{reference}.nii.gz" + ) + # Set flirt command + flirt_cmd = [ + "flirt", + "-interp spline", + f"-in {input['reference']['path'].resolve()}", + f"-ref {input['reference']['path'].resolve()}", + f"-applyisoxfm {input_resolution}", + f"-out {ref_path.resolve()}", + ] + # Call flirt + run_ext_cmd(name="flirt", cmd=flirt_cmd) + else: + logger.debug("Reference resolution matches input resolution") + ref_path = input["reference"]["path"] + + # Create a tempfile for warped output + applywarp_out_path = ( + element_tempdir + / f"warped_data_from_native_to_{reference}.nii.gz" + ) + # Set applywarp command + applywarp_cmd = [ + "applywarp", + "--interp=spline", + f"-i {input['path'].resolve()}", + # use resampled reference or original + f"-r {ref_path.resolve()}", + f"-w {warp_file_path.resolve()}", + f"-o {applywarp_out_path.resolve()}", + ] + # Call applywarp + run_ext_cmd(name="applywarp", cmd=applywarp_cmd) + + logger.debug("Updating warped data") + input.update( + { + # Update path to sync with "data" + "path": applywarp_out_path, + # Load nifti + "data": nib.load(applywarp_out_path), + # Switch space and prewarp_space + "space": reference, + "prewarp_space": input["space"], + # Save resampled reference path or overwrite original + # keeping it same + "reference": {"path": ref_path}, + } + ) + + # Check for data type's mask and warp if found + if input.get("mask") is not None: + logger.debug("Warping associated mask") + # Create a tempfile for warped mask output + applywarp_mask_out_path = ( + element_tempdir + / f"warped_mask_from_native_to_{reference}.nii.gz" + ) + # Set applywarp command + applywarp_mask_cmd = [ + "applywarp", + "--interp=nn", + f"-i {input['mask']['path'].resolve()}", + # use resampled reference or original + f"-r {ref_path.resolve()}", + f"-w {warp_file_path.resolve()}", + f"-o {applywarp_mask_out_path.resolve()}", + ] + # Call applywarp + run_ext_cmd(name="applywarp", cmd=applywarp_mask_cmd) + + logger.debug("Updating warped mask data") + input.update( + { + "mask": { + # Update path to sync with "data" + "path": applywarp_mask_out_path, + # Load nifti + "data": nib.load(applywarp_mask_out_path), + # Update mask's space + "space": reference, + } + } + ) + return input diff --git a/junifer/preprocess/warping/space_warper.py b/junifer/preprocess/warping/space_warper.py index 789ea9a63..a85e6955b 100644 --- a/junifer/preprocess/warping/space_warper.py +++ b/junifer/preprocess/warping/space_warper.py @@ -133,7 +133,7 @@ class SpaceWarper(BasePreprocessor): # Does not add any new keys return input_type - def preprocess( + def preprocess( # noqa: C901 self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, @@ -162,7 +162,7 @@ class SpaceWarper(BasePreprocessor): i.e., using ``"T1w"`` as reference. RuntimeError If warper could not be found in ``extra_input`` when - ``using="auto"`` or + ``using="auto"`` or converting from native space or if the data is in the correct space and does not require warping or if FSL is used when ``reference="T1w"``. @@ -184,6 +184,7 @@ class SpaceWarper(BasePreprocessor): input = FSLWarper().preprocess( input=input, extra_input=extra_input, + reference=self.reference, ) elif self.using == "ants": input = ANTsWarper().preprocess( @@ -204,6 +205,7 @@ class SpaceWarper(BasePreprocessor): input = FSLWarper().preprocess( input=input, extra_input=extra_input, + reference=self.reference, ) elif warper == "ants": input = ANTsWarper().preprocess( @@ -211,10 +213,11 @@ class SpaceWarper(BasePreprocessor): extra_input=extra_input, reference=self.reference, ) - # Transform to template space with ANTs possible - elif self.using == "ants" and self.reference != "T1w": + # Transform to template space + if self.using in ["fsl", "ants"] and self.reference != "T1w": + input_space = input["space"] # Check pre-requirements for space manipulation - if self.reference == input["space"]: + if self.using == "ants" and self.reference == input_space: raise_error( ( f"The target data is in {self.reference} space " @@ -224,20 +227,65 @@ class SpaceWarper(BasePreprocessor): ), klass=RuntimeError, ) - - input = ANTsWarper().preprocess( - input=input, - extra_input={}, - reference=self.reference, - ) - # Transform to template space with FSL not possible - elif self.using == "fsl" and self.reference != "T1w": - raise_error( - ( - f"Warping to {self.reference} space not possible with " - "FSL, use ANTs instead." - ), - klass=RuntimeError, - ) + # Transform from native to MNI possible conditionally + if input_space == "native": # pragma: no cover + # Check for reference as no T1w available + if input.get("reference") is None: + raise_error( + "`reference` key missing from input data type." + ) + # Check for extra inputs + if extra_input is None: + raise_error( + "No extra input provided, requires `Warp` " + "data type in particular." + ) + # Warp + input_prewarp_space = input["prewarp_space"] + warper = None + for entry in extra_input["Warp"]: + if ( + entry["src"] == input_space + and entry["dst"] == input_prewarp_space + ): + warper = entry["warper"] + if warper is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) + if warper == "fsl": + input = FSLWarper().preprocess( + input=input, + extra_input=extra_input, + reference=input_prewarp_space, + ) + elif warper == "ants": + input = ANTsWarper().preprocess( + input=input, + extra_input=extra_input, + reference=input_prewarp_space, + ) + else: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) + else: + # Transform from MNI to MNI template space not possible + if self.using == "fsl": + raise_error( + ( + f"Warping from {input_space} space to " + f"{self.reference} space not possible with " + "FSL, use ANTs instead." + ), + klass=RuntimeError, + ) + # Transform from MNI to MNI template space possible + else: + input = ANTsWarper().preprocess( + input=input, + extra_input={}, + reference=self.reference, + ) return input, None