[BUG]: Cannot use SpaceWarper to go from Native to template spaces #462

Merged
synchon merged 10 commits from update/bold-native-to-mni-warping-fsl into main 2025-09-24 12:48:10 +00:00
4 changed files with 345 additions and 121 deletions

View file

@ -0,0 +1 @@
Enable :class:`.SpaceWarper` to warp data from native space to template spaces via both FSL and ANTs by `Synchon Mandal`_

View file

@ -180,27 +180,74 @@ class ANTsWarper:
# Template space warping # Template space warping
else: else:
input_space = input["space"]
logger.debug( 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 # Native to MNI
xfm_file_path = get_xfm(src=input["space"], dst=reference) if input_space == "native":
# Get template space image # Get warp file path
template_space_img = get_template( xfm_file_path = None
space=reference, for entry in extra_input["Warp"]:
target_img=input["data"], if entry["src"] == "native" and entry["dst"] == reference:
extra_input=None, xfm_file_path = entry["path"]
) if xfm_file_path is None:
# Save template raise_error(
template_space_img_path = ( klass=RuntimeError,
element_tempdir / f"{reference}_T1w.nii.gz" msg="Could not find correct warp file path",
) )
nib.save(template_space_img, template_space_img_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 # Create a tempfile for warped output
warped_output_path = element_tempdir / ( 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 # Set antsApplyTransforms command
@ -210,7 +257,7 @@ class ANTsWarper:
"-e 3", "-e 3",
"-n LanczosWindowedSinc", "-n LanczosWindowedSinc",
f"-i {input['path'].resolve()}", f"-i {input['path'].resolve()}",
f"-r {template_space_img_path.resolve()}", f"-r {ref_path.resolve()}",
f"-t {xfm_file_path.resolve()}", f"-t {xfm_file_path.resolve()}",
f"-o {warped_output_path.resolve()}", f"-o {warped_output_path.resolve()}",
] ]
@ -226,18 +273,20 @@ class ANTsWarper:
"data": nib.load(warped_output_path), "data": nib.load(warped_output_path),
# Update warped input's space # Update warped input's space
"space": reference, "space": reference,
# Save reference path # Save resampled reference path or overwrite original
"reference": {"path": template_space_img_path}, # keeping it same
"reference": {"path": ref_path},
# Keep pre-warp space for further operations # 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 # Check for data type's mask and warp if found
if input.get("mask") is not None: if input.get("mask") is not None:
logger.debug("Warping associated mask")
# Create a tempfile for warped mask output # Create a tempfile for warped mask output
apply_transforms_mask_out_path = element_tempdir / ( 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 # Set antsApplyTransforms command
apply_transforms_mask_cmd = [ apply_transforms_mask_cmd = [
@ -246,8 +295,8 @@ class ANTsWarper:
"-e 3", "-e 3",
"-n 'GenericLabel[NearestNeighbor]'", "-n 'GenericLabel[NearestNeighbor]'",
f"-i {input['mask']['path'].resolve()}", f"-i {input['mask']['path'].resolve()}",
# use resampled reference # use resampled reference or original
f"-r {input['reference']['path'].resolve()}", f"-r {ref_path.resolve()}",
f"-t {xfm_file_path.resolve()}", f"-t {xfm_file_path.resolve()}",
f"-o {apply_transforms_mask_out_path.resolve()}", f"-o {apply_transforms_mask_out_path.resolve()}",
] ]

View file

@ -40,6 +40,7 @@ class FSLWarper:
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: dict[str, Any], extra_input: dict[str, Any],
reference: str,
) -> dict[str, Any]: # pragma: no cover ) -> dict[str, Any]: # pragma: no cover
"""Preprocess using FSL. """Preprocess using FSL.
@ -50,6 +51,10 @@ class FSLWarper:
extra_input : dict extra_input : dict
The other fields in the Junifer Data object. Should have ``T1w`` The other fields in the Junifer Data object. Should have ``T1w``
and ``Warp`` data types. 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 Returns
------- -------
@ -62,102 +67,223 @@ class FSLWarper:
If warp file path could not be found in ``extra_input``. 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 # Create element-specific tempdir for storing post-warping assets
element_tempdir = WorkDirManager().get_element_tempdir( element_tempdir = WorkDirManager().get_element_tempdir(
prefix="fsl_warper" prefix="fsl_warper"
) )
# Create a tempfile for resampled reference output # Warping to native space
flirt_out_path = element_tempdir / "resampled_reference.nii.gz" if reference == "T1w":
# Set flirt command logger.debug("Using FSL for space warping")
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 # Get the min of the voxel sizes from input and use it as the
applywarp_out_path = element_tempdir / "warped_data.nii.gz" # resolution
# Set applywarp command resolution = np.min(input["data"].header.get_zooms()[:3])
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)
logger.debug("Updating warped data") # Get warp file path
input.update( warp_file_path = None
{ for entry in extra_input["Warp"]:
# Update path to sync with "data" if entry["dst"] == "native":
"path": applywarp_out_path, warp_file_path = entry["path"]
# Load nifti if warp_file_path is None:
"data": nib.load(applywarp_out_path), raise_error(
# Use reference input's space as warped input's space klass=RuntimeError,
"space": extra_input["T1w"]["space"], msg="Could not find correct warp file path",
# 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 # Create a tempfile for resampled reference output
if input.get("mask") is not None: flirt_out_path = element_tempdir / "resampled_reference.nii.gz"
# Create a tempfile for warped mask output # Set flirt command
applywarp_mask_out_path = element_tempdir / "warped_mask.nii.gz" 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 # Set applywarp command
applywarp_mask_cmd = [ applywarp_cmd = [
"applywarp", "applywarp",
"--interp=nn", "--interp=spline",
f"-i {input['mask']['path'].resolve()}", f"-i {input['path'].resolve()}",
# use resampled reference # use resampled reference
f"-r {input['reference']['path'].resolve()}", f"-r {flirt_out_path.resolve()}",
f"-w {warp_file_path.resolve()}", f"-w {warp_file_path.resolve()}",
f"-o {applywarp_mask_out_path.resolve()}", f"-o {applywarp_out_path.resolve()}",
] ]
# Call applywarp # 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( input.update(
{ {
"mask": { # Update path to sync with "data"
# Update path to sync with "data" "path": applywarp_out_path,
"path": applywarp_mask_out_path, # Load nifti
# Load nifti "data": nib.load(applywarp_out_path),
"data": nib.load(applywarp_mask_out_path), # Use reference input's space as warped input's space
# Use reference input's space as warped input mask's "space": extra_input["T1w"]["space"],
# space # Save resampled reference path
"space": extra_input["T1w"]["space"], "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 return input

View file

@ -133,7 +133,7 @@ class SpaceWarper(BasePreprocessor):
# Does not add any new keys # Does not add any new keys
return input_type return input_type
def preprocess( def preprocess( # noqa: C901
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
@ -162,7 +162,7 @@ class SpaceWarper(BasePreprocessor):
i.e., using ``"T1w"`` as reference. i.e., using ``"T1w"`` as reference.
RuntimeError RuntimeError
If warper could not be found in ``extra_input`` when 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 if the data is in the correct space and does not require
warping or warping or
if FSL is used when ``reference="T1w"``. if FSL is used when ``reference="T1w"``.
@ -184,6 +184,7 @@ class SpaceWarper(BasePreprocessor):
input = FSLWarper().preprocess( input = FSLWarper().preprocess(
input=input, input=input,
extra_input=extra_input, extra_input=extra_input,
reference=self.reference,
) )
elif self.using == "ants": elif self.using == "ants":
input = ANTsWarper().preprocess( input = ANTsWarper().preprocess(
@ -204,6 +205,7 @@ class SpaceWarper(BasePreprocessor):
input = FSLWarper().preprocess( input = FSLWarper().preprocess(
input=input, input=input,
extra_input=extra_input, extra_input=extra_input,
reference=self.reference,
) )
elif warper == "ants": elif warper == "ants":
input = ANTsWarper().preprocess( input = ANTsWarper().preprocess(
@ -211,10 +213,11 @@ class SpaceWarper(BasePreprocessor):
extra_input=extra_input, extra_input=extra_input,
reference=self.reference, reference=self.reference,
) )
# Transform to template space with ANTs possible # Transform to template space
elif self.using == "ants" and self.reference != "T1w": if self.using in ["fsl", "ants"] and self.reference != "T1w":
input_space = input["space"]
# Check pre-requirements for space manipulation # Check pre-requirements for space manipulation
if self.reference == input["space"]: if self.using == "ants" and self.reference == input_space:
raise_error( raise_error(
( (
f"The target data is in {self.reference} space " f"The target data is in {self.reference} space "
@ -224,20 +227,65 @@ class SpaceWarper(BasePreprocessor):
), ),
klass=RuntimeError, klass=RuntimeError,
) )
# Transform from native to MNI possible conditionally
input = ANTsWarper().preprocess( if input_space == "native": # pragma: no cover
input=input, # Check for reference as no T1w available
extra_input={}, if input.get("reference") is None:
reference=self.reference, raise_error(
) "`reference` key missing from input data type."
# Transform to template space with FSL not possible )
elif self.using == "fsl" and self.reference != "T1w": # Check for extra inputs
raise_error( if extra_input is None:
( raise_error(
f"Warping to {self.reference} space not possible with " "No extra input provided, requires `Warp` "
"FSL, use ANTs instead." "data type in particular."
), )
klass=RuntimeError, # 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 return input, None