[ENH]: Update get_template interface #413

Merged
synchon merged 4 commits from update/get-template into main 2024-12-06 10:38:00 +00:00
7 changed files with 18 additions and 14 deletions

View file

@ -0,0 +1 @@
``target_data`` parameter of :func:`.get_template` has been replaced with ``target_img`` by `Fede Raimondo`_

View file

@ -125,7 +125,7 @@ class ANTsMaskWarper:
# Get template space image # Get template space image
template_space_img = get_template( template_space_img = get_template(
space=dst, space=dst,
target_data=target_data, target_img=mask_img,
extra_input=None, extra_input=None,
) )
# Save template to a tempfile # Save template to a tempfile

View file

@ -106,15 +106,15 @@ def compute_brain_mask(
if entry["dst"] == "native": if entry["dst"] == "native":
target_std_space = entry["src"] target_std_space = entry["src"]
target_img = target_data["data"]
# Fetch template in closest resolution # Fetch template in closest resolution
template = get_template( template = get_template(
space=target_std_space, space=target_std_space,
target_data=target_data, target_img=target_img,
extra_input=extra_input, extra_input=extra_input,
template_type=mask_type, template_type=mask_type,
) )
# Resample template to target image # Resample template to target image
target_img = target_data["data"]
resampled_template = resample_to_img( resampled_template = resample_to_img(
source_img=template, target_img=target_img source_img=template, target_img=target_img
) )

View file

@ -131,7 +131,7 @@ class ANTsParcellationWarper:
# Get template space image # Get template space image
template_space_img = get_template( template_space_img = get_template(
space=dst, space=dst,
target_data=target_data, target_img=parcellation_img,
extra_input=None, extra_input=None,
fraimondo commented 2024-12-05 18:19:06 +00:00 (Migrated from github.com)

i think this one needs to be in the refactor of the parcellation logic, otherwise it will fail tests here.

i think this one needs to be in the refactor of the parcellation logic, otherwise it will fail tests here.
synchon commented 2024-12-06 09:55:21 +00:00 (Migrated from github.com)

I updated it for both parcellation and mask to reduce the confusion in the other PRs.

I updated it for both parcellation and mask to reduce the confusion in the other PRs.
) )
# Save template to a tempfile # Save template to a tempfile

View file

@ -122,7 +122,7 @@ def get_xfm(
def get_template( def get_template(
space: str, space: str,
target_data: dict[str, Any], target_img: nib.Nifti1Image,
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
template_type: str = "T1w", template_type: str = "T1w",
) -> nib.Nifti1Image: ) -> nib.Nifti1Image:
@ -132,9 +132,9 @@ def get_template(
---------- ----------
space : str space : str
The name of the template space. The name of the template space.
target_data : dict target_img : Nifti1Image
The corresponding item of the data object for which the template space The corresponding image for which the template space will be loaded.
will be loaded. This is used to obtain the best matching resolution.
extra_input : dict, optional extra_input : dict, optional
The other fields in the data object. Useful for accessing other data The other fields in the data object. Useful for accessing other data
types (default None). types (default None).
@ -163,7 +163,6 @@ def get_template(
raise_error(f"Unknown template type: {template_type}") raise_error(f"Unknown template type: {template_type}")
# Get the min of the voxels sizes and use it as the resolution # Get the min of the voxels sizes and use it as the resolution
target_img = target_data["data"]
resolution = np.min(target_img.header.get_zooms()[:3]).astype(int) resolution = np.min(target_img.header.get_zooms()[:3]).astype(int)
# Fetch available resolutions for the template # Fetch available resolutions for the template

View file

@ -61,7 +61,9 @@ def test_get_template(template_type: str) -> None:
bold = element_data["BOLD"] bold = element_data["BOLD"]
# Get tailored parcellation # Get tailored parcellation
tailored_template = get_template( tailored_template = get_template(
space=bold["space"], target_data=bold, template_type=template_type space=bold["space"],
target_img=bold["data"],
template_type=template_type,
) )
assert isinstance(tailored_template, nib.Nifti1Image) assert isinstance(tailored_template, nib.Nifti1Image)
@ -74,7 +76,7 @@ def test_get_template_invalid_space() -> None:
vbm_gm = element_data["VBM_GM"] vbm_gm = element_data["VBM_GM"]
# Get tailored parcellation # Get tailored parcellation
with pytest.raises(ValueError, match="Unknown template space:"): with pytest.raises(ValueError, match="Unknown template space:"):
get_template(space="andromeda", target_data=vbm_gm) get_template(space="andromeda", target_img=vbm_gm["data"])
def test_get_template_invalid_template_type() -> None: def test_get_template_invalid_template_type() -> None:
@ -87,7 +89,7 @@ def test_get_template_invalid_template_type() -> None:
with pytest.raises(ValueError, match="Unknown template type:"): with pytest.raises(ValueError, match="Unknown template type:"):
get_template( get_template(
space=vbm_gm["space"], space=vbm_gm["space"],
target_data=vbm_gm, target_img=vbm_gm["data"],
template_type="xenon", template_type="xenon",
) )
@ -100,5 +102,7 @@ def test_get_template_closest_resolution() -> None:
vbm_gm = element_data["VBM_GM"] vbm_gm = element_data["VBM_GM"]
# Change header resolution to fetch closest resolution # Change header resolution to fetch closest resolution
element_data["VBM_GM"]["data"].header.set_zooms((3, 3, 3)) element_data["VBM_GM"]["data"].header.set_zooms((3, 3, 3))
template = get_template(space=vbm_gm["space"], target_data=vbm_gm) template = get_template(
space=vbm_gm["space"], target_img=vbm_gm["data"]
)
assert isinstance(template, nib.Nifti1Image) assert isinstance(template, nib.Nifti1Image)

View file

@ -189,7 +189,7 @@ class ANTsWarper:
# Get template space image # Get template space image
template_space_img = get_template( template_space_img = get_template(
space=reference, space=reference,
target_data=input, target_img=input["data"],
extra_input=None, extra_input=None,
) )
# Save template # Save template