[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
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"
)
# 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
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
template_space_img_path = (
element_tempdir / f"{reference}_T1w.nii.gz"
)
nib.save(template_space_img, template_space_img_path)
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()}",
]

View file

@ -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,6 +67,13 @@ class FSLWarper:
If warp file path could not be found in ``extra_input``.
"""
# Create element-specific tempdir for storing post-warping assets
element_tempdir = WorkDirManager().get_element_tempdir(
prefix="fsl_warper"
)
# Warping to native space
if reference == "T1w":
logger.debug("Using FSL for space warping")
# Get the min of the voxel sizes from input and use it as the
@ -75,12 +87,8 @@ class FSLWarper:
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"
klass=RuntimeError,
msg="Could not find correct warp file path",
)
# Create a tempfile for resampled reference output
@ -131,7 +139,9 @@ class FSLWarper:
# 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"
applywarp_mask_out_path = (
element_tempdir / "warped_mask.nii.gz"
)
# Set applywarp command
applywarp_mask_cmd = [
"applywarp",
@ -153,11 +163,127 @@ class FSLWarper:
"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
# 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

View file

@ -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,
)
# 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,
)
# 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,
)
return input, None