From 245b79791f9166010e0243b4fa5aaab9d9cb1e1b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:30:25 +0100 Subject: [PATCH 01/20] refactor: add support for list data type spec like Warp --- junifer/datagrabber/base.py | 7 +- junifer/datagrabber/datalad_base.py | 23 +- junifer/datagrabber/pattern.py | 99 ++++--- .../datagrabber/pattern_validation_mixin.py | 272 ++++++++++++------ junifer/pipeline/update_meta_mixin.py | 46 ++- 5 files changed, 308 insertions(+), 139 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 9e6d0246f..00b64cb69 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -96,7 +96,12 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): # Update metadata for _, t_val in out.items(): self.update_meta(t_val, "datagrabber") - t_val["meta"]["element"] = named_element + # Conditional for list dtype vals like Warp + if isinstance(t_val, list): + for entry in t_val: + entry["meta"]["element"] = named_element + else: + t_val["meta"]["element"] = named_element return out diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 2d4cb163b..55f3bb1fa 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -181,14 +181,21 @@ class DataladDataGrabber(BaseDataGrabber): """ to_get = [] for type_val in out.values(): - # Iterate to check for nested "types" like mask - for k, v in type_val.items(): - # Add base data type path - if k == "path": - to_get.append(v) - # Add nested data type path - if isinstance(v, dict) and "path" in v: - to_get.append(v["path"]) + # Conditional for list dtype vals like Warp + if isinstance(type_val, list): + for entry in type_val: + for k, v in entry.items(): + if k == "path": + to_get.append(v) + else: + # Iterate to check for nested "types" like mask + for k, v in type_val.items(): + # Add base data type path + if k == "path": + to_get.append(v) + # Add nested data type path + if isinstance(v, dict) and "path" in v: + to_get.append(v["path"]) if len(to_get) > 0: logger.debug(f"Getting {len(to_get)} files using datalad:") diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 9790e7dba..40588eaed 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -127,11 +127,21 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): "pattern": "...", "space": "...", }, - "Warp": { - "pattern": "...", - "src": "...", - "dst": "...", - } + } + + except ``Warp``, which needs to be a list of dictionaries as there can + be multiple spaces to warp (for example, with fMRIPrep): + + .. code-block:: none + + { + "Warp": [ + { + "pattern": "...", + "src": "...", + "dst": "...", + }, + ], } taken from :class:`.HCP1200`. @@ -380,42 +390,61 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): t_pattern = self.patterns[t_type] # Copy data type dictionary in output out[t_type] = deepcopy(t_pattern) - # Iterate to check for nested "types" like mask - for k, v in t_pattern.items(): - # Resolve pattern for base data type - if k == "pattern": - logger.info(f"Resolving path from pattern for {t_type}") - # Resolve pattern - base_data_type_pattern_path = self._get_path_from_patterns( - element=element, - pattern=v, - data_type=t_type, - ) - # Remove pattern key - out[t_type].pop("pattern") - # Add path key - out[t_type].update({"path": base_data_type_pattern_path}) - # Resolve pattern for nested data type - if isinstance(v, dict) and "pattern" in v: - # Set nested type key for easier access - t_nested_type = f"{t_type}.{k}" + # Conditional for list dtype vals like Warp + if isinstance(t_pattern, list): + for idx, entry in enumerate(t_pattern): logger.info( - f"Resolving path from pattern for {t_nested_type}" + f"Resolving path from pattern for {t_type}.{idx}" ) # Resolve pattern - nested_data_type_pattern_path = ( - self._get_path_from_patterns( - element=element, - pattern=v["pattern"], - data_type=t_nested_type, - ) + dtype_pattern_path = self._get_path_from_patterns( + element=element, + pattern=entry["pattern"], + data_type=f"{t_type}.{idx}", ) # Remove pattern key - out[t_type][k].pop("pattern") + out[t_type][idx].pop("pattern") # Add path key - out[t_type][k].update( - {"path": nested_data_type_pattern_path} - ) + out[t_type][idx].update({"path": dtype_pattern_path}) + else: + # Iterate to check for nested "types" like mask + for k, v in t_pattern.items(): + # Resolve pattern for base data type + if k == "pattern": + logger.info( + f"Resolving path from pattern for {t_type}" + ) + # Resolve pattern + base_dtype_pattern_path = self._get_path_from_patterns( + element=element, + pattern=v, + data_type=t_type, + ) + # Remove pattern key + out[t_type].pop("pattern") + # Add path key + out[t_type].update({"path": base_dtype_pattern_path}) + # Resolve pattern for nested data type + if isinstance(v, dict) and "pattern" in v: + # Set nested type key for easier access + t_nested_type = f"{t_type}.{k}" + logger.info( + f"Resolving path from pattern for {t_nested_type}" + ) + # Resolve pattern + nested_dtype_pattern_path = ( + self._get_path_from_patterns( + element=element, + pattern=v["pattern"], + data_type=t_nested_type, + ) + ) + # Remove pattern key + out[t_type][k].pop("pattern") + # Add path key + out[t_type][k].update( + {"path": nested_dtype_pattern_path} + ) return out diff --git a/junifer/datagrabber/pattern_validation_mixin.py b/junifer/datagrabber/pattern_validation_mixin.py index 9c130b902..90e63e2c1 100644 --- a/junifer/datagrabber/pattern_validation_mixin.py +++ b/junifer/datagrabber/pattern_validation_mixin.py @@ -3,7 +3,7 @@ # Authors: Synchon Mandal # License: AGPL -from typing import Dict, List +from typing import Dict, List, Union from ..utils import logger, raise_error, warn_with_log @@ -96,7 +96,7 @@ class PatternValidationMixin: def _validate_replacements( self, replacements: List[str], - patterns: Dict[str, Dict[str, str]], + patterns: Dict[str, Union[Dict[str, str], List[Dict[str, str]]]], partial_pattern_ok: bool, ) -> None: """Validate the replacements. @@ -136,14 +136,18 @@ class PatternValidationMixin: klass=TypeError, ) + # Make a list of all patterns recursively + all_patterns = [] + for dtype_val in patterns.values(): + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for entry in dtype_val: + all_patterns.append(entry.get("pattern", "")) + else: + all_patterns.append(dtype_val.get("pattern", "")) + # Check for stray replacements for x in replacements: - if all( - x not in y - for y in [ - data_type_val.get("pattern", "") - for data_type_val in patterns.values() - ] - ): + if all(x not in y for y in all_patterns): if partial_pattern_ok: warn_with_log( f"Replacement: `{x}` is not part of any pattern, " @@ -157,11 +161,19 @@ class PatternValidationMixin: # Check that at least one pattern has all the replacements at_least_one = False - for data_type_val in patterns.values(): - if all( - x in data_type_val.get("pattern", "") for x in replacements - ): - at_least_one = True + for dtype_val in patterns.values(): + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for entry in dtype_val: + if all( + x in entry.get("pattern", "") for x in replacements + ): + at_least_one = True + else: + if all( + x in dtype_val.get("pattern", "") for x in replacements + ): + at_least_one = True if not at_least_one and not partial_pattern_ok: raise_error( msg="At least one pattern must contain all replacements." @@ -251,7 +263,7 @@ class PatternValidationMixin: self, types: List[str], replacements: List[str], - patterns: Dict[str, Dict[str, str]], + patterns: Dict[str, Union[Dict[str, str], List[Dict[str, str]]]], partial_pattern_ok: bool = False, ) -> None: """Validate the patterns. @@ -298,87 +310,185 @@ class PatternValidationMixin: msg="`patterns` must contain all `types`", klass=ValueError ) # Check against schema - for data_type_key, data_type_val in patterns.items(): + for dtype_key, dtype_val in patterns.items(): # Check if valid data type is provided - if data_type_key not in PATTERNS_SCHEMA: + if dtype_key not in PATTERNS_SCHEMA: raise_error( - f"Unknown data type: {data_type_key}, " + f"Unknown data type: {dtype_key}, " f"should be one of: {list(PATTERNS_SCHEMA.keys())}" ) - # Check mandatory keys for data type - self._validate_mandatory_keys( - keys=list(data_type_val), - schema=PATTERNS_SCHEMA[data_type_key]["mandatory"], - data_type=data_type_key, - partial_pattern_ok=partial_pattern_ok, - ) - # Check optional keys for data type - for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][ - "optional" - ].items(): - if optional_key not in data_type_val: - logger.debug( - f"Optional key: `{optional_key}` missing for " - f"{data_type_key}" - ) - else: - logger.debug( - f"Optional key: `{optional_key}` found for " - f"{data_type_key}" - ) - # Set nested type name for easier access - nested_data_type = f"{data_type_key}.{optional_key}" - nested_mandatory_keys_schema = PATTERNS_SCHEMA[ - data_type_key - ]["optional"][optional_key]["mandatory"] - nested_optional_keys_schema = PATTERNS_SCHEMA[ - data_type_key - ]["optional"][optional_key]["optional"] - # Check mandatory keys for nested type + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for idx, entry in enumerate(dtype_val): + # Check mandatory keys for data type self._validate_mandatory_keys( - keys=list(optional_val["mandatory"]), - schema=nested_mandatory_keys_schema, - data_type=nested_data_type, + keys=list(entry), + schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], + data_type=f"{dtype_key}.{idx}", partial_pattern_ok=partial_pattern_ok, ) - # Check optional keys for nested type - for nested_optional_key in nested_optional_keys_schema: - if nested_optional_key not in optional_val["optional"]: + # Check optional keys for data type + for optional_key, optional_val in PATTERNS_SCHEMA[ + dtype_key + ]["optional"].items(): + if optional_key not in entry: logger.debug( - f"Optional key: `{nested_optional_key}` " - f"missing for {nested_data_type}" + f"Optional key: `{optional_key}` missing for " + f"{dtype_key}.{idx}" ) else: logger.debug( - f"Optional key: `{nested_optional_key}` found " - f"for {nested_data_type}" + f"Optional key: `{optional_key}` found for " + f"{dtype_key}.{idx}" ) - # Check stray key for nested data type + # Set nested type name for easier access + nested_dtype = f"{dtype_key}.{idx}.{optional_key}" + nested_mandatory_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["mandatory"] + nested_optional_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["optional"] + # Check mandatory keys for nested type + self._validate_mandatory_keys( + keys=list(optional_val["mandatory"]), + schema=nested_mandatory_keys_schema, + data_type=nested_dtype, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for nested type + for ( + nested_optional_key + ) in nested_optional_keys_schema: + if ( + nested_optional_key + not in optional_val["optional"] + ): + logger.debug( + f"Optional key: " + f"`{nested_optional_key}` missing for " + f"{nested_dtype}" + ) + else: + logger.debug( + f"Optional key: " + f"`{nested_optional_key}` found for " + f"{nested_dtype}" + ) + # Check stray key for nested data type + self._identify_stray_keys( + keys=( + optional_val["mandatory"] + + optional_val["optional"] + ), + schema=( + nested_mandatory_keys_schema + + nested_optional_keys_schema + ), + data_type=nested_dtype, + ) + # Check stray key for data type self._identify_stray_keys( - keys=optional_val["mandatory"] - + optional_val["optional"], - schema=nested_mandatory_keys_schema - + nested_optional_keys_schema, - data_type=nested_data_type, + keys=list(entry.keys()), + schema=( + PATTERNS_SCHEMA[dtype_key]["mandatory"] + + list( + PATTERNS_SCHEMA[dtype_key]["optional"].keys() + ) + ), + data_type=dtype_key, ) - # Check stray key for data type - self._identify_stray_keys( - keys=list(data_type_val.keys()), - schema=( - PATTERNS_SCHEMA[data_type_key]["mandatory"] - + list(PATTERNS_SCHEMA[data_type_key]["optional"].keys()) - ), - data_type=data_type_key, - ) - # Wildcard check in patterns - if "}*" in data_type_val.get("pattern", ""): - raise_error( - msg=( - f"`{data_type_key}.pattern` must not contain `*` " - "following a replacement" - ), - klass=ValueError, + # Wildcard check in patterns + if "}*" in entry.get("pattern", ""): + raise_error( + msg=( + f"`{dtype_key}.pattern` must not contain `*` " + "following a replacement" + ), + klass=ValueError, + ) + else: + # Check mandatory keys for data type + self._validate_mandatory_keys( + keys=list(dtype_val), + schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], + data_type=dtype_key, + partial_pattern_ok=partial_pattern_ok, ) + # Check optional keys for data type + for optional_key, optional_val in PATTERNS_SCHEMA[dtype_key][ + "optional" + ].items(): + if optional_key not in dtype_val: + logger.debug( + f"Optional key: `{optional_key}` missing for " + f"{dtype_key}." + ) + else: + logger.debug( + f"Optional key: `{optional_key}` found for " + f"{dtype_key}." + ) + # Set nested type name for easier access + nested_dtype = f"{dtype_key}.{optional_key}" + nested_mandatory_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["mandatory"] + nested_optional_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["optional"] + # Check mandatory keys for nested type + self._validate_mandatory_keys( + keys=list(optional_val["mandatory"]), + schema=nested_mandatory_keys_schema, + data_type=nested_dtype, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for nested type + for nested_optional_key in nested_optional_keys_schema: + if ( + nested_optional_key + not in optional_val["optional"] + ): + logger.debug( + f"Optional key: `{nested_optional_key}` " + f"missing for {nested_dtype}" + ) + else: + logger.debug( + f"Optional key: `{nested_optional_key}` " + f"found for {nested_dtype}" + ) + # Check stray key for nested data type + self._identify_stray_keys( + keys=( + optional_val["mandatory"] + + optional_val["optional"] + ), + schema=( + nested_mandatory_keys_schema + + nested_optional_keys_schema + ), + data_type=nested_dtype, + ) + # Check stray key for data type + self._identify_stray_keys( + keys=list(dtype_val.keys()), + schema=( + PATTERNS_SCHEMA[dtype_key]["mandatory"] + + list(PATTERNS_SCHEMA[dtype_key]["optional"].keys()) + ), + data_type=dtype_key, + ) + # Wildcard check in patterns + if "}*" in dtype_val.get("pattern", ""): + raise_error( + msg=( + f"`{dtype_key}.pattern` must not contain `*` " + "following a replacement" + ), + klass=ValueError, + ) # Validate replacements self._validate_replacements( diff --git a/junifer/pipeline/update_meta_mixin.py b/junifer/pipeline/update_meta_mixin.py index c88978244..68a910778 100644 --- a/junifer/pipeline/update_meta_mixin.py +++ b/junifer/pipeline/update_meta_mixin.py @@ -36,17 +36,35 @@ class UpdateMetaMixin: for k, v in vars(self).items(): if not k.startswith("_"): t_meta[k] = v - # Add "meta" to the step's local context dict - if "meta" not in input: - input["meta"] = {} - # Add step name - input["meta"][step_name] = t_meta - # Add step dependencies - if "dependencies" not in input["meta"]: - input["meta"]["dependencies"] = set() - # Update step dependencies - dependencies = getattr(self, "_DEPENDENCIES", set()) - if dependencies is not None: - if not isinstance(dependencies, (set, list)): - dependencies = {dependencies} - input["meta"]["dependencies"].update(dependencies) + # Conditional for list dtype vals like Warp + if isinstance(input, list): + for entry in input: + # Add "meta" to the step's entry's local context dict + if "meta" not in entry: + entry["meta"] = {} + # Add step name + entry["meta"][step_name] = t_meta + # Add step dependencies + if "dependencies" not in entry["meta"]: + entry["meta"]["dependencies"] = set() + # Update step dependencies + dependencies = getattr(self, "_DEPENDENCIES", set()) + if dependencies is not None: + if not isinstance(dependencies, (set, list)): + dependencies = {dependencies} + entry["meta"]["dependencies"].update(dependencies) + else: + # Add "meta" to the step's local context dict + if "meta" not in input: + input["meta"] = {} + # Add step name + input["meta"][step_name] = t_meta + # Add step dependencies + if "dependencies" not in input["meta"]: + input["meta"]["dependencies"] = set() + # Update step dependencies + dependencies = getattr(self, "_DEPENDENCIES", set()) + if dependencies is not None: + if not isinstance(dependencies, (set, list)): + dependencies = {dependencies} + input["meta"]["dependencies"].update(dependencies) -- 2.52.0 From 4e7b1e6e9e5ad291750c4891b09b4ce4412cdc6b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:33:20 +0100 Subject: [PATCH 02/20] refactor: adapt Warp data types for datagrabbers --- junifer/datagrabber/aomic/id1000.py | 29 +++++++++++++++++-------- junifer/datagrabber/aomic/piop1.py | 29 +++++++++++++++++-------- junifer/datagrabber/aomic/piop2.py | 29 +++++++++++++++++-------- junifer/datagrabber/dmcc13_benchmark.py | 29 +++++++++++++++++-------- junifer/datagrabber/hcp1200/hcp1200.py | 23 ++++++++++++++------ 5 files changed, 96 insertions(+), 43 deletions(-) diff --git a/junifer/datagrabber/aomic/id1000.py b/junifer/datagrabber/aomic/id1000.py index 7d0586255..d5b930f9f 100644 --- a/junifer/datagrabber/aomic/id1000.py +++ b/junifer/datagrabber/aomic/id1000.py @@ -169,15 +169,26 @@ class DataladAOMICID1000(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/aomic/piop1.py b/junifer/datagrabber/aomic/piop1.py index d2763bb18..fe2e9e643 100644 --- a/junifer/datagrabber/aomic/piop1.py +++ b/junifer/datagrabber/aomic/piop1.py @@ -204,15 +204,26 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/aomic/piop2.py b/junifer/datagrabber/aomic/piop2.py index a9852ab75..ef58ba4c8 100644 --- a/junifer/datagrabber/aomic/piop2.py +++ b/junifer/datagrabber/aomic/piop2.py @@ -202,15 +202,26 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/dmcc13_benchmark.py b/junifer/datagrabber/dmcc13_benchmark.py index af4d32f96..5ea6ae8ac 100644 --- a/junifer/datagrabber/dmcc13_benchmark.py +++ b/junifer/datagrabber/dmcc13_benchmark.py @@ -221,15 +221,26 @@ class DMCC13Benchmark(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep-1.3.2/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep-1.3.2/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + }, + { + "pattern": ( + "derivatives/fmriprep-1.3.2/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/hcp1200/hcp1200.py b/junifer/datagrabber/hcp1200/hcp1200.py index 0fa8a8018..1d956aab8 100644 --- a/junifer/datagrabber/hcp1200/hcp1200.py +++ b/junifer/datagrabber/hcp1200/hcp1200.py @@ -122,13 +122,22 @@ class HCP1200(PatternDataGrabber): "pattern": "{subject}/T1w/T1w_acpc_dc_restore.nii.gz", "space": "native", }, - "Warp": { - "pattern": ( - "{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz" - ), - "src": "MNI152NLin6Asym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz" + ), + "src": "MNI152NLin6Asym", + "dst": "native", + }, + { + "pattern": ( + "{subject}/MNINonLinear/xfms/acpc_dc2standard.nii.gz" + ), + "src": "native", + "dst": "MNI152NLin6Asym", + }, + ], } # The replacements replacements = ["subject", "task", "phase_encoding"] -- 2.52.0 From 2a69a24c970d0f15e1d45c314a287bbfc942283d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:35:06 +0100 Subject: [PATCH 03/20] update: add warper parameter to Warp entries --- junifer/datagrabber/aomic/id1000.py | 2 ++ junifer/datagrabber/aomic/piop1.py | 2 ++ junifer/datagrabber/aomic/piop2.py | 2 ++ junifer/datagrabber/dmcc13_benchmark.py | 2 ++ junifer/datagrabber/hcp1200/hcp1200.py | 2 ++ junifer/datagrabber/pattern.py | 3 ++- junifer/datagrabber/pattern_validation_mixin.py | 2 +- 7 files changed, 13 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/aomic/id1000.py b/junifer/datagrabber/aomic/id1000.py index d5b930f9f..cf5fc90e4 100644 --- a/junifer/datagrabber/aomic/id1000.py +++ b/junifer/datagrabber/aomic/id1000.py @@ -178,6 +178,7 @@ class DataladAOMICID1000(PatternDataladDataGrabber): ), "src": "MNI152NLin2009cAsym", "dst": "native", + "warper": "ants", }, { "pattern": ( @@ -187,6 +188,7 @@ class DataladAOMICID1000(PatternDataladDataGrabber): ), "src": "native", "dst": "MNI152NLin2009cAsym", + "warper": "ants", }, ], } diff --git a/junifer/datagrabber/aomic/piop1.py b/junifer/datagrabber/aomic/piop1.py index fe2e9e643..29751e842 100644 --- a/junifer/datagrabber/aomic/piop1.py +++ b/junifer/datagrabber/aomic/piop1.py @@ -213,6 +213,7 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): ), "src": "MNI152NLin2009cAsym", "dst": "native", + "warper": "ants", }, { "pattern": ( @@ -222,6 +223,7 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): ), "src": "native", "dst": "MNI152NLin2009cAsym", + "warper": "ants", }, ], } diff --git a/junifer/datagrabber/aomic/piop2.py b/junifer/datagrabber/aomic/piop2.py index ef58ba4c8..53d922d13 100644 --- a/junifer/datagrabber/aomic/piop2.py +++ b/junifer/datagrabber/aomic/piop2.py @@ -211,6 +211,7 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber): ), "src": "MNI152NLin2009cAsym", "dst": "native", + "warper": "ants", }, { "pattern": ( @@ -220,6 +221,7 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber): ), "src": "native", "dst": "MNI152NLin2009cAsym", + "warper": "ants", }, ], } diff --git a/junifer/datagrabber/dmcc13_benchmark.py b/junifer/datagrabber/dmcc13_benchmark.py index 5ea6ae8ac..173e10887 100644 --- a/junifer/datagrabber/dmcc13_benchmark.py +++ b/junifer/datagrabber/dmcc13_benchmark.py @@ -230,6 +230,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber): ), "src": "MNI152NLin2009cAsym", "dst": "native", + "warper": "ants", }, { "pattern": ( @@ -239,6 +240,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber): ), "src": "native", "dst": "MNI152NLin2009cAsym", + "warper": "ants", }, ], } diff --git a/junifer/datagrabber/hcp1200/hcp1200.py b/junifer/datagrabber/hcp1200/hcp1200.py index 1d956aab8..c2455d16e 100644 --- a/junifer/datagrabber/hcp1200/hcp1200.py +++ b/junifer/datagrabber/hcp1200/hcp1200.py @@ -129,6 +129,7 @@ class HCP1200(PatternDataGrabber): ), "src": "MNI152NLin6Asym", "dst": "native", + "warper": "fsl", }, { "pattern": ( @@ -136,6 +137,7 @@ class HCP1200(PatternDataGrabber): ), "src": "native", "dst": "MNI152NLin6Asym", + "warper": "fsl", }, ], } diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 40588eaed..25012b074 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -89,7 +89,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): .. code-block:: none { - "mandatory": ["pattern", "src", "dst"], + "mandatory": ["pattern", "src", "dst", "warper"], "optional": [] } @@ -140,6 +140,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): "pattern": "...", "src": "...", "dst": "...", + "warper": "...", }, ], } diff --git a/junifer/datagrabber/pattern_validation_mixin.py b/junifer/datagrabber/pattern_validation_mixin.py index 90e63e2c1..fcae6c8e7 100644 --- a/junifer/datagrabber/pattern_validation_mixin.py +++ b/junifer/datagrabber/pattern_validation_mixin.py @@ -36,7 +36,7 @@ PATTERNS_SCHEMA = { }, }, "Warp": { - "mandatory": ["pattern", "src", "dst"], + "mandatory": ["pattern", "src", "dst", "warper"], "optional": {}, }, "VBM_GM": { -- 2.52.0 From c04ee308f4c3376484a410caf2503804cda22856 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:52:31 +0100 Subject: [PATCH 04/20] refactor: adapt warpers to use new Warp spec --- .../coordinates/_ants_coordinates_warper.py | 22 +++++++++--- junifer/data/coordinates/_coordinates.py | 25 +++++++------- .../coordinates/_fsl_coordinates_warper.py | 22 +++++++++--- junifer/data/masks/_ants_mask_warper.py | 20 +++++++++-- junifer/data/masks/_fsl_mask_warper.py | 22 +++++++++--- junifer/data/masks/_masks.py | 34 +++++++++---------- .../_ants_parcellation_warper.py | 20 +++++++++-- .../parcellations/_fsl_parcellation_warper.py | 23 ++++++++++--- junifer/data/parcellations/_parcellations.py | 30 ++++++++-------- junifer/preprocess/warping/_ants_warper.py | 20 +++++++++-- junifer/preprocess/warping/_fsl_warper.py | 19 +++++++++-- 11 files changed, 187 insertions(+), 70 deletions(-) diff --git a/junifer/data/coordinates/_ants_coordinates_warper.py b/junifer/data/coordinates/_ants_coordinates_warper.py index e1d7f35db..5852e6f60 100644 --- a/junifer/data/coordinates/_ants_coordinates_warper.py +++ b/junifer/data/coordinates/_ants_coordinates_warper.py @@ -9,7 +9,7 @@ import numpy as np from numpy.typing import ArrayLike from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["ANTsCoordinatesWarper"] @@ -39,17 +39,31 @@ class ANTsCoordinatesWarper: will be applied. extra_input : dict, optional The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates - (default None). + data kinds that needs to be used in the computation of coordinates. Returns ------- numpy.ndarray The transformed coordinates. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using ANTs for coordinates transformation") + # 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="ants_coordinates_warper" @@ -79,7 +93,7 @@ class ANTsCoordinatesWarper: "-f 0", f"-i {pretransform_coordinates_path.resolve()}", f"-o {transformed_coords_path.resolve()}", - f"-t {extra_input['Warp']['path'].resolve()};", + f"-t {warp_file_path.resolve()}", ] # Call antsApplyTransformsToPoints run_ext_cmd( diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index 01073de9b..9b08b8125 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -311,7 +311,7 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". + If warper could not be found in ``extra_input``. ValueError If ``extra_input`` is None when ``target_data``'s space is native. @@ -329,27 +329,26 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): f"{target_data['space']} space for further computation." ) - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # Check for warper to use correct tool + warper = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + warper = entry["warper"] + if warper is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) + if warper == "fsl": seeds = FSLCoordinatesWarper().warp( seeds=seeds, target_data=target_data, extra_input=extra_input, ) - elif warp_file_ext == ".h5": + elif warper == "ants": seeds = ANTsCoordinatesWarper().warp( seeds=seeds, target_data=target_data, extra_input=extra_input, ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, - ) return seeds, labels diff --git a/junifer/data/coordinates/_fsl_coordinates_warper.py b/junifer/data/coordinates/_fsl_coordinates_warper.py index 72948e0f9..614a7638d 100644 --- a/junifer/data/coordinates/_fsl_coordinates_warper.py +++ b/junifer/data/coordinates/_fsl_coordinates_warper.py @@ -9,7 +9,7 @@ import numpy as np from numpy.typing import ArrayLike from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["FSLCoordinatesWarper"] @@ -39,17 +39,31 @@ class FSLCoordinatesWarper: will be applied. extra_input : dict, optional The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates - (default None). + data kinds that needs to be used in the computation of coordinates. Returns ------- numpy.ndarray The transformed coordinates. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using FSL for coordinates transformation") + # 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_coordinates_warper" @@ -72,7 +86,7 @@ class FSLCoordinatesWarper: "| img2imgcoord -mm", f"-src {target_data['path'].resolve()}", f"-dest {target_data['reference_path'].resolve()}", - f"-warp {extra_input['Warp']['path'].resolve()}", + f"-warp {warp_file_path.resolve()}", f"> {transformed_coords_path.resolve()};", f"sed -i 1d {transformed_coords_path.resolve()}", ] diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index 572c63097..94b691084 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd from ..template_spaces import get_template, get_xfm @@ -66,6 +66,11 @@ class ANTsMaskWarper: nibabel.nifti1.Nifti1Image The transformed mask image. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ # Create element-scoped tempdir so that warped mask is # available later as nibabel stores file path reference for @@ -83,6 +88,17 @@ class ANTsMaskWarper: if dst == "T1w": logger.debug("Using ANTs for mask transformation") + # 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", + ) + # Save existing mask image to a tempfile prewarp_mask_path = element_tempdir / "prewarp_mask.nii.gz" nib.save(mask_img, prewarp_mask_path) @@ -98,7 +114,7 @@ class ANTsMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_file_path.resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/masks/_fsl_mask_warper.py b/junifer/data/masks/_fsl_mask_warper.py index 933d49b90..7215a1ce6 100644 --- a/junifer/data/masks/_fsl_mask_warper.py +++ b/junifer/data/masks/_fsl_mask_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd if TYPE_CHECKING: @@ -46,17 +46,31 @@ class FSLMaskWarper: will be applied. extra_input : dict, optional The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of mask - (default None). + data kinds that needs to be used in the computation of mask. Returns ------- nibabel.nifti1.Nifti1Image The transformed mask image. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using FSL for mask transformation") + # 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-scoped tempdir so that warped mask is # available later as nibabel stores file path reference for # loading on computation @@ -77,7 +91,7 @@ class FSLMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_file_path.resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index be7a1b688..fcde59c2c 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -105,7 +105,9 @@ def compute_brain_mask( "data type to infer target template space." ) # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + target_std_space = entry["src"] # Fetch template in closest resolution template = get_template( @@ -358,7 +360,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". + If warper could not be found in ``extra_input``. ValueError If extra key is provided in addition to mask name in ``masks`` or if no mask is provided or @@ -383,8 +385,16 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) - # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + # Set target standard space to warp file space source and warper + warper = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + target_std_space = entry["src"] + warper = entry["warper"] + if warper is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -510,17 +520,15 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp mask if target data is native if target_space == "native": - # extra_input check done earlier - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # extra_input check done earlier and warper exists + if warper == "fsl": mask_img = FSLMaskWarper().warp( mask_name="native", mask_img=mask_img, target_data=target_data, extra_input=extra_input, ) - elif warp_file_ext == ".h5": + elif warper == "ants": mask_img = ANTsMaskWarper().warp( mask_name="native", mask_img=mask_img, @@ -529,14 +537,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): target_data=target_data, extra_input=extra_input, ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, - ) return mask_img diff --git a/junifer/data/parcellations/_ants_parcellation_warper.py b/junifer/data/parcellations/_ants_parcellation_warper.py index 64de58ef0..e1a47f499 100644 --- a/junifer/data/parcellations/_ants_parcellation_warper.py +++ b/junifer/data/parcellations/_ants_parcellation_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd from ..template_spaces import get_template, get_xfm @@ -66,6 +66,11 @@ class ANTsParcellationWarper: nibabel.nifti1.Nifti1Image The transformed parcellation image. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ # Create element-scoped tempdir so that warped parcellation is # available later as nibabel stores file path reference for @@ -83,6 +88,17 @@ class ANTsParcellationWarper: if dst == "T1w": logger.debug("Using ANTs for parcellation transformation") + # 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", + ) + # Save existing parcellation image to a tempfile prewarp_parcellation_path = ( element_tempdir / "prewarp_parcellation.nii.gz" @@ -102,7 +118,7 @@ class ANTsParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_file_path.resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/parcellations/_fsl_parcellation_warper.py b/junifer/data/parcellations/_fsl_parcellation_warper.py index 7fc9717ca..0b1b54e75 100644 --- a/junifer/data/parcellations/_fsl_parcellation_warper.py +++ b/junifer/data/parcellations/_fsl_parcellation_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd if TYPE_CHECKING: @@ -46,17 +46,32 @@ class FSLParcellationWarper: will be applied. extra_input : dict, optional The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of parcellation - (default None). + data kinds that needs to be used in the computation of + parcellation. Returns ------- nibabel.nifti1.Nifti1Image The transformed parcellation image. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using FSL for parcellation transformation") + # 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-scoped tempdir so that warped parcellation is # available later as nibabel stores file path reference for # loading on computation @@ -81,7 +96,7 @@ class FSLParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_file_path.resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 3e9cb4195..d59e8f68a 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -395,7 +395,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". + If warper could not be found in ``extra_input``. ValueError If ``extra_input`` is None when ``target_data``'s space is native. @@ -413,8 +413,16 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) - # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + # Set target standard space to warp file space source and warper + warper = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + target_std_space = entry["src"] + warper = entry["warper"] + if warper is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -471,17 +479,15 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp parcellation if target space is native if target_space == "native": - # extra_input check done earlier - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # extra_input check done earlier and warper exists + if warper == "fsl": resampled_parcellation_img = FSLParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, target_data=target_data, extra_input=extra_input, ) - elif warp_file_ext == ".h5": + elif warper == "ants": resampled_parcellation_img = ANTsParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, @@ -490,14 +496,6 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): target_data=target_data, extra_input=extra_input, ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, - ) return resampled_parcellation_img, labels diff --git a/junifer/preprocess/warping/_ants_warper.py b/junifer/preprocess/warping/_ants_warper.py index dfc4e7bf8..17c49ae6b 100644 --- a/junifer/preprocess/warping/_ants_warper.py +++ b/junifer/preprocess/warping/_ants_warper.py @@ -15,7 +15,7 @@ import numpy as np from ...data import get_template, get_xfm from ...pipeline import WorkDirManager from ...typing import Dependencies, ExternalDependencies -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["ANTsWarper"] @@ -63,6 +63,11 @@ class ANTsWarper: values and new ``reference_path`` key whose value points to the reference file used for warping. + Raises + ------ + RuntimeError + 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( @@ -77,6 +82,17 @@ class ANTsWarper: # 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 a tempfile for resampled reference output resample_image_out_path = ( element_tempdir / "resampled_reference.nii.gz" @@ -105,7 +121,7 @@ class ANTsWarper: f"-i {input['path'].resolve()}", # use resampled reference f"-r {resample_image_out_path.resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_file_path.resolve()}", f"-o {apply_transforms_out_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/preprocess/warping/_fsl_warper.py b/junifer/preprocess/warping/_fsl_warper.py index c04ae6582..2e912df93 100644 --- a/junifer/preprocess/warping/_fsl_warper.py +++ b/junifer/preprocess/warping/_fsl_warper.py @@ -14,7 +14,7 @@ import numpy as np from ...pipeline import WorkDirManager from ...typing import Dependencies, ExternalDependencies -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["FSLWarper"] @@ -59,6 +59,11 @@ class FSLWarper: values and new ``reference_path`` key whose value points to the reference file used for warping. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using FSL for space warping") @@ -66,6 +71,16 @@ class FSLWarper: # 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" @@ -93,7 +108,7 @@ class FSLWarper: "--interp=spline", f"-i {input['path'].resolve()}", f"-r {flirt_out_path.resolve()}", # use resampled reference - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_file_path.resolve()}", f"-o {applywarp_out_path.resolve()}", ] # Call applywarp -- 2.52.0 From 561a4864cea6d8279c827017afcc75516ad8f732 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:54:23 +0100 Subject: [PATCH 05/20] refactor: allow for "auto" in conditional dependencies --- junifer/pipeline/pipeline_step_mixin.py | 16 +++++++--- junifer/preprocess/warping/space_warper.py | 35 +++++++++++++++++++--- 2 files changed, 43 insertions(+), 8 deletions(-) diff --git a/junifer/pipeline/pipeline_step_mixin.py b/junifer/pipeline/pipeline_step_mixin.py index b4f45594f..3e1c12892 100644 --- a/junifer/pipeline/pipeline_step_mixin.py +++ b/junifer/pipeline/pipeline_step_mixin.py @@ -196,10 +196,18 @@ class PipelineStepMixin: for dependency in obj._CONDITIONAL_DEPENDENCIES: if dependency["using"] == obj.using: depends_on = dependency["depends_on"] - # Check dependencies - _check_dependencies(depends_on) - # Check external dependencies - _check_ext_dependencies(depends_on) + # Conditional to make `using="auto"` work + if isinstance(depends_on, list): + for entry in depends_on: + # Check dependencies + _check_dependencies(entry) + # Check external dependencies + _check_ext_dependencies(entry) + else: + # Check dependencies + _check_dependencies(depends_on) + # Check external dependencies + _check_ext_dependencies(depends_on) # Check dependencies _check_dependencies(self) diff --git a/junifer/preprocess/warping/space_warper.py b/junifer/preprocess/warping/space_warper.py index 445bc4c5c..fdab703c3 100644 --- a/junifer/preprocess/warping/space_warper.py +++ b/junifer/preprocess/warping/space_warper.py @@ -24,11 +24,12 @@ class SpaceWarper(BasePreprocessor): Parameters ---------- - using : {"fsl", "ants"} + using : {"fsl", "ants", "auto"} Implementation to use for warping: * "fsl" : Use FSL's ``applywarp`` * "ants" : Use ANTs' ``antsApplyTransforms`` + * "auto" : Auto-select tool when ``reference="T1w"`` reference : str The data type to use as reference for warping, can be either a data @@ -56,6 +57,10 @@ class SpaceWarper(BasePreprocessor): "using": "ants", "depends_on": ANTsWarper, }, + { + "using": "auto", + "depends_on": [FSLWarper, ANTsWarper], + }, ] def __init__( @@ -156,14 +161,16 @@ class SpaceWarper(BasePreprocessor): If ``extra_input`` is None when transforming to native space i.e., using ``"T1w"`` as reference. RuntimeError - If the data is in the correct space and does not require + If warper could not be found in ``extra_input`` when + ``using="auto"`` or + if the data is in the correct space and does not require warping or - if FSL is used for template space warping. + if FSL is used when ``reference="T1w"``. """ logger.info(f"Warping to {self.reference} space using SpaceWarper") # Transform to native space - if self.using in ["fsl", "ants"] and self.reference == "T1w": + if self.using in ["fsl", "ants", "auto"] and self.reference == "T1w": # Check for extra inputs if extra_input is None: raise_error( @@ -182,6 +189,26 @@ class SpaceWarper(BasePreprocessor): extra_input=extra_input, reference=self.reference, ) + elif self.using == "auto": + warper = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + 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, + ) + elif warper == "ants": + input = ANTsWarper().preprocess( + input=input, + extra_input=extra_input, + reference=self.reference, + ) # Transform to template space with ANTs possible elif self.using == "ants" and self.reference != "T1w": # Check pre-requirements for space manipulation -- 2.52.0 From 25948630e7d8008b3c68473c8c9fe73ead933415 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:58:07 +0100 Subject: [PATCH 06/20] chore: fix mask pattern in DMCC13Benchmark --- junifer/datagrabber/dmcc13_benchmark.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/datagrabber/dmcc13_benchmark.py b/junifer/datagrabber/dmcc13_benchmark.py index 173e10887..7111624b1 100644 --- a/junifer/datagrabber/dmcc13_benchmark.py +++ b/junifer/datagrabber/dmcc13_benchmark.py @@ -150,7 +150,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber): "mask": { "pattern": ( "derivatives/fmriprep-1.3.2/{subject}/{session}/" - "/func/{subject}_{session}_task-{task}_acq-mb4" + "func/{subject}_{session}_task-{task}_acq-mb4" "{phase_encoding}_run-{run}_" "space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz" ), -- 2.52.0 From b4a43eedb378e41e493de4b3885ae9043da289ad Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:58:46 +0100 Subject: [PATCH 07/20] docs: update extending/dependencies.rst --- docs/extending/dependencies.rst | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/docs/extending/dependencies.rst b/docs/extending/dependencies.rst index 441ba5414..48132f990 100644 --- a/docs/extending/dependencies.rst +++ b/docs/extending/dependencies.rst @@ -39,7 +39,7 @@ and others who use it will thank you. Handling external dependencies from toolboxes --------------------------------------------- -You can also specify dependencies of external toolboxes like AFIN, FSL and ANTs, +You can also specify dependencies of external toolboxes like AFNI, FSL and ANTs, by having a class attribute like so: .. code-block:: python @@ -87,6 +87,10 @@ that it shows the problem a bit better and how we solve it: "using": "ants", "depends_on": ANTsWarper, }, + { + "using": "auto", + "depends_on": [FSLWarper, ANTsWarper], + }, ] def __init__( @@ -100,13 +104,20 @@ Here, you see a new class attribute ``_CONDITIONAL_DEPENDENCIES`` which is a list of dictionaries with two keys: * ``using`` (str) : lowercased name of the toolbox -* ``depends_on`` (object) : a class which implements the particular tool's use +* ``depends_on`` (object or list of objects) : a class or list of classes which \ + implements the particular tool's use It is mandatory to have the ``using`` positional argument in the constructor in this case as the validation starts with this and moves further. It is also mandatory to only allow the value of ``using`` argument to be one of them specified in the ``using`` key of ``_CONDITIONAL_DEPENDENCIES`` entries. +For some cases, like we have here, it might be worth to have ``"auto"`` for +``"using"`` which eases the choice of a particular tool by the user and +instead lets ``junifer`` do it automatically. In that case, ``"depends_on"`` +needs to be a list of the tool implementation classes. This also requires the +user to have all the tools in the ``PATH``. + For brevity, we only show the ``FSLWarper`` here but ``ANTsWarper`` looks very similar. ``FSLWarper`` looks like this (only the relevant part is shown here): -- 2.52.0 From 2958a2bda2748843f48d39b06cd2a6da8baf3613 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 11:59:05 +0100 Subject: [PATCH 08/20] docs: update understanding/preprocess.rst --- docs/understanding/preprocess.rst | 28 ++++++++++++++++++++-------- 1 file changed, 20 insertions(+), 8 deletions(-) diff --git a/docs/understanding/preprocess.rst b/docs/understanding/preprocess.rst index c5701aa18..c9b7946ab 100644 --- a/docs/understanding/preprocess.rst +++ b/docs/understanding/preprocess.rst @@ -146,14 +146,24 @@ Warping to subject's native space To warp to subject's native space, the dataset needs to provide ``T1w`` and ``Warp`` data types and the DataGrabber needs to at least have ``["BOLD", "T1w", "Warp"]`` (if you are warping ``BOLD``) as the ``types`` -parameter's value. The :class:`.SpaceWarper`'s ``reference`` parameter needs +parameter's value. + +The :class:`.SpaceWarper`'s ``reference`` parameter needs to be set to ``T1w``, which means that the ``BOLD`` data will be transformed using the ``T1w`` as reference (it's resampled internally to match the -resolution of the ``BOLD``). The ``Warp`` data type is new and it's only purpose -is to provide the warp or transformation file (can be linear, non-linear or -linear + non-linear transform) for the purpose. For ``using`` parameter, you can -pass either ``"fsl"`` or ``"ants"`` depending on the warp or transformation file -format. +resolution of the ``BOLD``). + +The ``Warp`` data type provides the warp or transformation file (can be linear, +non-linear or linear + non-linear transform) for the purpose. For ``using`` +parameter, you can pass either ``fsl`` or ``ants`` depending on the warp or +transformation file format. You can also provide ``auto`` to ``using`` in which +case either ``FSL`` or ``ANTs`` will be used based on the preferred warping tool +or the one available as provided by the DataGrabber. This also requires that +both the tools are in the ``PATH``. + +And finally, you would need to set the ``on`` parameter to ``BOLD`` to make it +clear which data type you intend to warp, as the :class:`.SpaceWarper` is also +capable of warping ``T1w``. An example YAML might look like this: @@ -163,6 +173,7 @@ An example YAML might look like this: - kind: SpaceWarper using: fsl reference: T1w + on: BOLD .. _preprocess_warping_template: @@ -174,8 +185,8 @@ data type that you want to work on) in ``MNI152NLin6Asym`` template space but you would like to compute features in ``MNI152NLin2009cAsym`` template space, you can also use the :class:`.SpaceWarper` by setting the ``reference`` parameter to the template space's name, in this case, -``reference="MNI152NLin2009cAsym"``. The ``using`` parameter needs to be set -to ``"ants"`` as we need it to warp the data. +``reference: MNI152NLin2009cAsym``. The ``using`` parameter needs to be set +to ``ants`` as we need it to warp the data. .. note:: @@ -190,3 +201,4 @@ For an YAML example: - kind: SpaceWarper using: ants reference: MNI152NLin2009cAsym + on: BOLD -- 2.52.0 From 219562bae0415a3758aef002dad311dd61c7f4af Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 12:00:07 +0100 Subject: [PATCH 09/20] chore: update log messages for PatternValidationMixin --- junifer/datagrabber/pattern_validation_mixin.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/junifer/datagrabber/pattern_validation_mixin.py b/junifer/datagrabber/pattern_validation_mixin.py index fcae6c8e7..e9e1a8602 100644 --- a/junifer/datagrabber/pattern_validation_mixin.py +++ b/junifer/datagrabber/pattern_validation_mixin.py @@ -132,7 +132,7 @@ class PatternValidationMixin: if any(not isinstance(x, str) for x in replacements): raise_error( - msg="`replacements` must be a list of strings.", + msg="`replacements` must be a list of strings", klass=TypeError, ) @@ -152,11 +152,11 @@ class PatternValidationMixin: warn_with_log( f"Replacement: `{x}` is not part of any pattern, " "things might not work as expected if you are unsure " - "of what you are doing" + "of what you are doing." ) else: raise_error( - msg=f"Replacement: {x} is not part of any pattern." + msg=f"Replacement: `{x}` is not part of any pattern" ) # Check that at least one pattern has all the replacements @@ -176,7 +176,7 @@ class PatternValidationMixin: at_least_one = True if not at_least_one and not partial_pattern_ok: raise_error( - msg="At least one pattern must contain all replacements." + msg="At least one pattern must contain all replacements" ) def _validate_mandatory_keys( @@ -219,7 +219,7 @@ class PatternValidationMixin: warn_with_log( f"Mandatory key: `{key}` not found for {data_type}, " "things might not work as expected if you are unsure " - "of what you are doing" + "of what you are doing." ) else: raise_error( @@ -227,7 +227,7 @@ class PatternValidationMixin: klass=KeyError, ) else: - logger.debug(f"Mandatory key: `{key}` found for {data_type}") + logger.debug(f"Mandatory key: `{key}` found for {data_type}.") def _identify_stray_keys( self, keys: List[str], schema: List[str], data_type: str -- 2.52.0 From 1c7cc974898ecc661f1101083e34cc1f9872a018 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 15:57:34 +0100 Subject: [PATCH 10/20] chore: update mock dataset creation scripts for new Warp spec --- tools/create_dmcc13_benchmark_example_dataset.py | 6 +++++- tools/create_hcp1200_example_dataset.py | 14 +++++++++----- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/tools/create_dmcc13_benchmark_example_dataset.py b/tools/create_dmcc13_benchmark_example_dataset.py index 8be13edb3..996a897c3 100644 --- a/tools/create_dmcc13_benchmark_example_dataset.py +++ b/tools/create_dmcc13_benchmark_example_dataset.py @@ -78,7 +78,11 @@ if __name__ == "__main__": ( f"anat/sub-{sub:02d}_from-MNI152NLin2009cAsym_to-T1w_" "mode-image_xfm.h5" - ), # Warp + ), # Warp to native + ( + f"anat/sub-{sub:02d}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), # Warp from native ] # Tasks for functional data diff --git a/tools/create_hcp1200_example_dataset.py b/tools/create_hcp1200_example_dataset.py index 0759cbb9b..9929bbd9f 100644 --- a/tools/create_hcp1200_example_dataset.py +++ b/tools/create_hcp1200_example_dataset.py @@ -44,11 +44,15 @@ if __name__ == "__main__": sub_warp_datadir = subdir / "MNINonLinear" / "xfms" # Create subject data directory sub_warp_datadir.mkdir(parents=True) - # Set subject data file - sub_warp_datafile = sub_warp_datadir / "standard2acpc_dc.nii.gz" - # Write subject data file - with open(sub_warp_datafile, "w") as f: - f.write("placeholder") + # Set subject data files + sub_warp_datafiles = [ + sub_warp_datadir / "standard2acpc_dc.nii.gz", + sub_warp_datadir / "acpc_dc2standard.nii.gz", + ] + # Write subject data files + for datafile in sub_warp_datafiles: + with open(datafile, "w") as f: + f.write("placeholder") # BOLD data for task, phase_encoding in product( -- 2.52.0 From abc2676530dbfb9f4b1d3ae5ae69d8f23fed7e37 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Nov 2024 15:58:06 +0100 Subject: [PATCH 11/20] chore: update tests for DMCC13Benchmark --- .../tests/test_dmcc13_benchmark.py | 35 ++++++++++++++----- 1 file changed, 26 insertions(+), 9 deletions(-) diff --git a/junifer/datagrabber/tests/test_dmcc13_benchmark.py b/junifer/datagrabber/tests/test_dmcc13_benchmark.py index 2b2566de1..ca41950bb 100644 --- a/junifer/datagrabber/tests/test_dmcc13_benchmark.py +++ b/junifer/datagrabber/tests/test_dmcc13_benchmark.py @@ -116,7 +116,12 @@ def test_DMCC13Benchmark( data_file_names.extend( [ "sub-01_desc-preproc_T1w.nii.gz", - "sub-01_from-MNI152NLin2009cAsym_to-T1w_mode-image_xfm.h5", + [ + "sub-01_from-MNI152NLin2009cAsym_to-T1w" + "_mode-image_xfm.h5", + "sub-01_from-T1w_to-MNI152NLin2009cAsym" + "_mode-image_xfm.h5", + ], ] ) else: @@ -127,14 +132,26 @@ def test_DMCC13Benchmark( for data_type, data_file_name in zip(data_types, data_file_names): # Assert data type assert data_type in out - # Assert data file path exists - assert out[data_type]["path"].exists() - # Assert data file path is a file - assert out[data_type]["path"].is_file() - # Assert data file name - assert out[data_type]["path"].name == data_file_name - # Assert metadata - assert "meta" in out[data_type] + # Conditional for Warp + if data_type == "Warp": + for idx, fname in enumerate(data_file_name): + # Assert data file path exists + assert out[data_type][idx]["path"].exists() + # Assert data file path is a file + assert out[data_type][idx]["path"].is_file() + # Assert data file name + assert out[data_type][idx]["path"].name == fname + # Assert metadata + assert "meta" in out[data_type][idx] + else: + # Assert data file path exists + assert out[data_type]["path"].exists() + # Assert data file path is a file + assert out[data_type]["path"].is_file() + # Assert data file name + assert out[data_type]["path"].name == data_file_name + # Assert metadata + assert "meta" in out[data_type] # Check BOLD nested data types for type_, file_name in zip( -- 2.52.0 From 3397d1f52b77b48cbe23cffd4cbc2a706ab02da6 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 12 Nov 2024 15:52:32 +0100 Subject: [PATCH 12/20] docs: improve understanding/preprocess.rst --- docs/understanding/preprocess.rst | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/docs/understanding/preprocess.rst b/docs/understanding/preprocess.rst index c9b7946ab..08cd73ddb 100644 --- a/docs/understanding/preprocess.rst +++ b/docs/understanding/preprocess.rst @@ -157,9 +157,8 @@ The ``Warp`` data type provides the warp or transformation file (can be linear, non-linear or linear + non-linear transform) for the purpose. For ``using`` parameter, you can pass either ``fsl`` or ``ants`` depending on the warp or transformation file format. You can also provide ``auto`` to ``using`` in which -case either ``FSL`` or ``ANTs`` will be used based on the preferred warping tool -or the one available as provided by the DataGrabber. This also requires that -both the tools are in the ``PATH``. +case either ``FSL`` or ``ANTs`` will be used based on the file format provided +by the DataGrabber. This also requires that both the tools are in the ``PATH``. And finally, you would need to set the ``on`` parameter to ``BOLD`` to make it clear which data type you intend to warp, as the :class:`.SpaceWarper` is also -- 2.52.0 From a0ac6256e2ac5e01910fd19cabf045a3c7b2d3a0 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 12 Nov 2024 16:02:53 +0100 Subject: [PATCH 13/20] refactor: remove duplicated logic for PipelineStepMixin deps check --- junifer/pipeline/pipeline_step_mixin.py | 14 +++++--------- 1 file changed, 5 insertions(+), 9 deletions(-) diff --git a/junifer/pipeline/pipeline_step_mixin.py b/junifer/pipeline/pipeline_step_mixin.py index 3e1c12892..84590dc2e 100644 --- a/junifer/pipeline/pipeline_step_mixin.py +++ b/junifer/pipeline/pipeline_step_mixin.py @@ -197,17 +197,13 @@ class PipelineStepMixin: if dependency["using"] == obj.using: depends_on = dependency["depends_on"] # Conditional to make `using="auto"` work - if isinstance(depends_on, list): - for entry in depends_on: - # Check dependencies - _check_dependencies(entry) - # Check external dependencies - _check_ext_dependencies(entry) - else: + if not isinstance(depends_on, list): + depends_on = [depends_on] + for entry in depends_on: # Check dependencies - _check_dependencies(depends_on) + _check_dependencies(entry) # Check external dependencies - _check_ext_dependencies(depends_on) + _check_ext_dependencies(entry) # Check dependencies _check_dependencies(self) -- 2.52.0 From 81b2ed5cdb429f0b378b16ffa912cef77782f487 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 12 Nov 2024 16:03:27 +0100 Subject: [PATCH 14/20] refactor: remove duplicated code for UpdateMetaMixin meta update --- junifer/pipeline/update_meta_mixin.py | 40 +++++++++------------------ 1 file changed, 13 insertions(+), 27 deletions(-) diff --git a/junifer/pipeline/update_meta_mixin.py b/junifer/pipeline/update_meta_mixin.py index 68a910778..abbbc07d0 100644 --- a/junifer/pipeline/update_meta_mixin.py +++ b/junifer/pipeline/update_meta_mixin.py @@ -4,7 +4,7 @@ # Synchon Mandal # License: AGPL -from typing import Dict +from typing import Dict, List, Union __all__ = ["UpdateMetaMixin"] @@ -15,14 +15,14 @@ class UpdateMetaMixin: def update_meta( self, - input: Dict, + input: Union[Dict, List[Dict]], step_name: str, ) -> None: """Update metadata. Parameters ---------- - input : dict + input : dict or list of dict The data object to update. step_name : str The name of the pipeline step. @@ -37,34 +37,20 @@ class UpdateMetaMixin: if not k.startswith("_"): t_meta[k] = v # Conditional for list dtype vals like Warp - if isinstance(input, list): - for entry in input: - # Add "meta" to the step's entry's local context dict - if "meta" not in entry: - entry["meta"] = {} - # Add step name - entry["meta"][step_name] = t_meta - # Add step dependencies - if "dependencies" not in entry["meta"]: - entry["meta"]["dependencies"] = set() - # Update step dependencies - dependencies = getattr(self, "_DEPENDENCIES", set()) - if dependencies is not None: - if not isinstance(dependencies, (set, list)): - dependencies = {dependencies} - entry["meta"]["dependencies"].update(dependencies) - else: - # Add "meta" to the step's local context dict - if "meta" not in input: - input["meta"] = {} + if not isinstance(input, list): + input = [input] + for entry in input: + # Add "meta" to the step's entry's local context dict + if "meta" not in entry: + entry["meta"] = {} # Add step name - input["meta"][step_name] = t_meta + entry["meta"][step_name] = t_meta # Add step dependencies - if "dependencies" not in input["meta"]: - input["meta"]["dependencies"] = set() + if "dependencies" not in entry["meta"]: + entry["meta"]["dependencies"] = set() # Update step dependencies dependencies = getattr(self, "_DEPENDENCIES", set()) if dependencies is not None: if not isinstance(dependencies, (set, list)): dependencies = {dependencies} - input["meta"]["dependencies"].update(dependencies) + entry["meta"]["dependencies"].update(dependencies) -- 2.52.0 From 0e21e453a73807e012b573473d784a0ad69a65a1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 10:56:02 +0100 Subject: [PATCH 15/20] refactor: update the interface and logic for data warpers --- .../coordinates/_ants_coordinates_warper.py | 26 +++------------- junifer/data/coordinates/_coordinates.py | 4 +-- .../coordinates/_fsl_coordinates_warper.py | 26 +++------------- junifer/data/masks/_ants_mask_warper.py | 29 ++++++----------- junifer/data/masks/_fsl_mask_warper.py | 26 +++------------- junifer/data/masks/_masks.py | 6 ++-- .../_ants_parcellation_warper.py | 31 +++++++------------ .../parcellations/_fsl_parcellation_warper.py | 27 +++------------- junifer/data/parcellations/_parcellations.py | 6 ++-- 9 files changed, 49 insertions(+), 132 deletions(-) diff --git a/junifer/data/coordinates/_ants_coordinates_warper.py b/junifer/data/coordinates/_ants_coordinates_warper.py index 5852e6f60..d9b4ccf09 100644 --- a/junifer/data/coordinates/_ants_coordinates_warper.py +++ b/junifer/data/coordinates/_ants_coordinates_warper.py @@ -9,7 +9,7 @@ import numpy as np from numpy.typing import ArrayLike from ...pipeline import WorkDirManager -from ...utils import logger, raise_error, run_ext_cmd +from ...utils import logger, run_ext_cmd __all__ = ["ANTsCoordinatesWarper"] @@ -26,7 +26,7 @@ class ANTsCoordinatesWarper: self, seeds: ArrayLike, target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> ArrayLike: """Warp ``seeds`` to correct space. @@ -37,33 +37,17 @@ class ANTsCoordinatesWarper: target_data : dict The corresponding item of the data object to which the coordinates will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates. + warp_data : dict or None + The warp data item of the data object. Returns ------- numpy.ndarray The transformed coordinates. - Raises - ------ - RuntimeError - If warp file path could not be found in ``extra_input``. - """ logger.debug("Using ANTs for coordinates transformation") - # 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="ants_coordinates_warper" @@ -93,7 +77,7 @@ class ANTsCoordinatesWarper: "-f 0", f"-i {pretransform_coordinates_path.resolve()}", f"-o {transformed_coords_path.resolve()}", - f"-t {warp_file_path.resolve()}", + f"-t {warp_data['path'].resolve()}", ] # Call antsApplyTransformsToPoints run_ext_cmd( diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index 9b08b8125..ad66d6dd9 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -342,13 +342,13 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): seeds = FSLCoordinatesWarper().warp( seeds=seeds, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) elif warper == "ants": seeds = ANTsCoordinatesWarper().warp( seeds=seeds, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) return seeds, labels diff --git a/junifer/data/coordinates/_fsl_coordinates_warper.py b/junifer/data/coordinates/_fsl_coordinates_warper.py index 614a7638d..afe65280c 100644 --- a/junifer/data/coordinates/_fsl_coordinates_warper.py +++ b/junifer/data/coordinates/_fsl_coordinates_warper.py @@ -9,7 +9,7 @@ import numpy as np from numpy.typing import ArrayLike from ...pipeline import WorkDirManager -from ...utils import logger, raise_error, run_ext_cmd +from ...utils import logger, run_ext_cmd __all__ = ["FSLCoordinatesWarper"] @@ -26,7 +26,7 @@ class FSLCoordinatesWarper: self, seeds: ArrayLike, target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> ArrayLike: """Warp ``seeds`` to correct space. @@ -37,33 +37,17 @@ class FSLCoordinatesWarper: target_data : dict The corresponding item of the data object to which the coordinates will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates. + warp_data : dict + The warp data item of the data object. Returns ------- numpy.ndarray The transformed coordinates. - Raises - ------ - RuntimeError - If warp file path could not be found in ``extra_input``. - """ logger.debug("Using FSL for coordinates transformation") - # 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_coordinates_warper" @@ -86,7 +70,7 @@ class FSLCoordinatesWarper: "| img2imgcoord -mm", f"-src {target_data['path'].resolve()}", f"-dest {target_data['reference_path'].resolve()}", - f"-warp {warp_file_path.resolve()}", + f"-warp {warp_data['path'].resolve()}", f"> {transformed_coords_path.resolve()};", f"sed -i 1d {transformed_coords_path.resolve()}", ] diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index 94b691084..23acd916c 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -34,7 +34,7 @@ class ANTsMaskWarper: src: str, dst: str, target_data: Dict[str, Any], - extra_input: Optional[Dict[str, Any]] = None, + warp_data: Optional[Dict[str, Any]], ) -> "Nifti1Image": """Warp ``mask_img`` to correct space. @@ -55,11 +55,9 @@ class ANTsMaskWarper: target_data : dict The corresponding item of the data object to which the mask will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of mask - (default None). - + warp_data : dict or None + The warp data item of the data object. The value is unused if + ``dst!="T1w"``. Returns ------- @@ -69,7 +67,7 @@ class ANTsMaskWarper: Raises ------ RuntimeError - If warp file path could not be found in ``extra_input``. + If ``warp_data`` is None when ``dst="T1w"``. """ # Create element-scoped tempdir so that warped mask is @@ -86,18 +84,11 @@ class ANTsMaskWarper: # Native space warping if dst == "T1w": - logger.debug("Using ANTs for mask transformation") + # Warp data check + if warp_data is None: + raise_error("No `warp_data` provided") - # 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", - ) + logger.debug("Using ANTs for mask transformation") # Save existing mask image to a tempfile prewarp_mask_path = element_tempdir / "prewarp_mask.nii.gz" @@ -114,7 +105,7 @@ class ANTsMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {warp_file_path.resolve()}", + f"-t {warp_data['path'].resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/masks/_fsl_mask_warper.py b/junifer/data/masks/_fsl_mask_warper.py index 7215a1ce6..4aab3ed02 100644 --- a/junifer/data/masks/_fsl_mask_warper.py +++ b/junifer/data/masks/_fsl_mask_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, raise_error, run_ext_cmd +from ...utils import logger, run_ext_cmd if TYPE_CHECKING: @@ -31,7 +31,7 @@ class FSLMaskWarper: mask_name: str, mask_img: "Nifti1Image", target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> "Nifti1Image": """Warp ``mask_img`` to correct space. @@ -44,33 +44,17 @@ class FSLMaskWarper: target_data : dict The corresponding item of the data object to which the mask will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of mask. + warp_data : dict + The warp data item of the data object. Returns ------- nibabel.nifti1.Nifti1Image The transformed mask image. - Raises - ------ - RuntimeError - If warp file path could not be found in ``extra_input``. - """ logger.debug("Using FSL for mask transformation") - # 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-scoped tempdir so that warped mask is # available later as nibabel stores file path reference for # loading on computation @@ -91,7 +75,7 @@ class FSLMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {warp_file_path.resolve()}", + f"-w {warp_data['path'].resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index fcde59c2c..478efba56 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -499,7 +499,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): src=mask_space, dst=target_std_space, target_data=target_data, - extra_input=None, + warp_data=None, ) all_masks.append(mask_img) @@ -526,7 +526,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): mask_name="native", mask_img=mask_img, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) elif warper == "ants": mask_img = ANTsMaskWarper().warp( @@ -535,7 +535,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): src="", dst="T1w", target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) return mask_img diff --git a/junifer/data/parcellations/_ants_parcellation_warper.py b/junifer/data/parcellations/_ants_parcellation_warper.py index e1a47f499..4d86d8c78 100644 --- a/junifer/data/parcellations/_ants_parcellation_warper.py +++ b/junifer/data/parcellations/_ants_parcellation_warper.py @@ -34,7 +34,7 @@ class ANTsParcellationWarper: src: str, dst: str, target_data: Dict[str, Any], - extra_input: Optional[Dict[str, Any]] = None, + warp_data: Optional[Dict[str, Any]], ) -> "Nifti1Image": """Warp ``parcellation_img`` to correct space. @@ -55,11 +55,9 @@ class ANTsParcellationWarper: target_data : dict The corresponding item of the data object to which the parcellation will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of parcellation - (default None). - + warp_data : dict or None + The warp data item of the data object. The value is unused if + ``dst!="T1w"``. Returns ------- @@ -68,8 +66,8 @@ class ANTsParcellationWarper: Raises ------ - RuntimeError - If warp file path could not be found in ``extra_input``. + ValueError + If ``warp_data`` is None when ``dst="T1w"``. """ # Create element-scoped tempdir so that warped parcellation is @@ -86,18 +84,11 @@ class ANTsParcellationWarper: # Native space warping if dst == "T1w": - logger.debug("Using ANTs for parcellation transformation") + # Warp data check + if warp_data is None: + raise_error("No `warp_data` provided") - # 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", - ) + logger.debug("Using ANTs for parcellation transformation") # Save existing parcellation image to a tempfile prewarp_parcellation_path = ( @@ -118,7 +109,7 @@ class ANTsParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {warp_file_path.resolve()}", + f"-t {warp_data['path'].resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/parcellations/_fsl_parcellation_warper.py b/junifer/data/parcellations/_fsl_parcellation_warper.py index 0b1b54e75..f9af85504 100644 --- a/junifer/data/parcellations/_fsl_parcellation_warper.py +++ b/junifer/data/parcellations/_fsl_parcellation_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, raise_error, run_ext_cmd +from ...utils import logger, run_ext_cmd if TYPE_CHECKING: @@ -31,7 +31,7 @@ class FSLParcellationWarper: parcellation_name: str, parcellation_img: "Nifti1Image", target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> "Nifti1Image": """Warp ``parcellation_img`` to correct space. @@ -44,34 +44,17 @@ class FSLParcellationWarper: target_data : dict The corresponding item of the data object to which the parcellation will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of - parcellation. + warp_data : dict + The warp data item of the data object. Returns ------- nibabel.nifti1.Nifti1Image The transformed parcellation image. - Raises - ------ - RuntimeError - If warp file path could not be found in ``extra_input``. - """ logger.debug("Using FSL for parcellation transformation") - # 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-scoped tempdir so that warped parcellation is # available later as nibabel stores file path reference for # loading on computation @@ -96,7 +79,7 @@ class FSLParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {warp_file_path.resolve()}", + f"-w {warp_data['path'].resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index d59e8f68a..b8439ddb3 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -449,7 +449,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): src=space, dst=target_std_space, target_data=target_data, - extra_input=None, + warp_data=None, ) # Remove extra dimension added by ANTs img = image.math_img("np.squeeze(img)", img=raw_img) @@ -485,7 +485,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): parcellation_name="native", parcellation_img=resampled_parcellation_img, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) elif warper == "ants": resampled_parcellation_img = ANTsParcellationWarper().warp( @@ -494,7 +494,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): src="", dst="T1w", target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) return resampled_parcellation_img, labels -- 2.52.0 From be9efa1041125eabb98065405990f9a5dc84624e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 10:57:06 +0100 Subject: [PATCH 16/20] update: add utility function to get native space warper and adjust space warpers --- junifer/data/utils.py | 70 +++++++++++++++++++++- junifer/preprocess/warping/_ants_warper.py | 5 ++ junifer/preprocess/warping/_fsl_warper.py | 3 +- 3 files changed, 74 insertions(+), 4 deletions(-) diff --git a/junifer/data/utils.py b/junifer/data/utils.py index dd73a5388..a03171417 100644 --- a/junifer/data/utils.py +++ b/junifer/data/utils.py @@ -4,14 +4,14 @@ # Synchon Mandal # License: AGPL -from typing import List, Optional, Union +from typing import Dict, List, MutableMapping, Optional, Union import numpy as np -from ..utils.logging import logger +from ..utils import logger, raise_error -__all__ = ["closest_resolution"] +__all__ = ["closest_resolution", "get_native_warper"] def closest_resolution( @@ -49,3 +49,67 @@ def closest_resolution( closest = np.min(valid_resolution) return closest + + +def get_native_warper( + target_data: MutableMapping, + other_data: MutableMapping, + inverse: bool = False, +) -> Dict: + """Get correct warping specification for native space. + + Parameters + ---------- + target_data : dict + The target data from the pipeline data object. + other_data : dict + The other data in the pipeline data object. + inverse : bool, optional + Whether to get the inverse warping specification (default False). + + Returns + ------- + dict + The correct warping specification. + + Raises + ------ + RuntimeError + If no warper or multiple possible warpers are found. + + """ + # Get possible warpers + possible_warpers = [] + for entry in other_data["Warp"]: + if not inverse: + if ( + entry["src"] == target_data["prewarp_space"] + and entry["dst"] == "native" + ): + possible_warpers.append(entry) + else: + if ( + entry["dst"] == target_data["prewarp_space"] + and entry["src"] == "native" + ): + possible_warpers.append(entry) + + # Check for no warper + if not possible_warpers: + raise_error( + klass=RuntimeError, + msg="Could not find correct warping specification", + ) + + # Check for multiple possible warpers + if len(possible_warpers) > 1: + raise_error( + klass=RuntimeError, + msg=( + "Cannot proceed as multiple warping specification found, " + "adjust either the DataGrabber or the working space: " + f"{possible_warpers}" + ), + ) + + return possible_warpers[0] diff --git a/junifer/preprocess/warping/_ants_warper.py b/junifer/preprocess/warping/_ants_warper.py index 17c49ae6b..9fdcaa0e6 100644 --- a/junifer/preprocess/warping/_ants_warper.py +++ b/junifer/preprocess/warping/_ants_warper.py @@ -131,6 +131,8 @@ class ANTsWarper: input["data"] = nib.load(apply_transforms_out_path) # Save resampled reference path input["reference_path"] = resample_image_out_path + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] # Use reference input's space as warped input's space input["space"] = extra_input["T1w"]["space"] @@ -179,6 +181,9 @@ class ANTsWarper: # Modify target data input["data"] = nib.load(warped_output_path) + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] + # Update warped input's space input["space"] = reference return input diff --git a/junifer/preprocess/warping/_fsl_warper.py b/junifer/preprocess/warping/_fsl_warper.py index 2e912df93..1261ba4a6 100644 --- a/junifer/preprocess/warping/_fsl_warper.py +++ b/junifer/preprocess/warping/_fsl_warper.py @@ -118,7 +118,8 @@ class FSLWarper: input["data"] = nib.load(applywarp_out_path) # Save resampled reference path input["reference_path"] = flirt_out_path - + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] # Use reference input's space as warped input's space input["space"] = extra_input["T1w"]["space"] -- 2.52.0 From c0a338e73f1242dc8599fc68fba7b2c9ce418ccc Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 10:59:18 +0100 Subject: [PATCH 17/20] refactor: use get_native_warper() in tailored data fetch --- junifer/data/coordinates/_coordinates.py | 38 +++++++++++++------- junifer/data/masks/_masks.py | 32 ++++++++--------- junifer/data/parcellations/_parcellations.py | 36 +++++++++---------- 3 files changed, 57 insertions(+), 49 deletions(-) diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index ad66d6dd9..e51b9fb9d 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -14,6 +14,7 @@ from numpy.typing import ArrayLike from ...utils import logger, raise_error from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry +from ..utils import get_native_warper from ._ants_coordinates_warper import ANTsCoordinatesWarper from ._fsl_coordinates_warper import FSLCoordinatesWarper @@ -311,7 +312,8 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ RuntimeError - If warper could not be found in ``extra_input``. + If warping specification required for warping using ANTs, is not + found. ValueError If ``extra_input`` is None when ``target_data``'s space is native. @@ -329,22 +331,34 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): f"{target_data['space']} space for further computation." ) - # Check for warper to use correct tool - warper = None - for entry in extra_input["Warp"]: - if entry["dst"] == "native": - warper = entry["warper"] - if warper is None: - raise_error( - klass=RuntimeError, msg="Could not find correct warper" - ) - if warper == "fsl": + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) + # Conditional for warping tool implementation + if warper_spec["warper"] == "fsl": seeds = FSLCoordinatesWarper().warp( seeds=seeds, target_data=target_data, warp_data=warper_spec, ) - elif warper == "ants": + elif warper_spec["warper"] == "ants": + # Requires the inverse warp + inverse_warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + inverse=True, + ) + # Check warper + if inverse_warper_spec["warper"] != "ants": + raise_error( + klass=RuntimeError, + msg=( + "Warping specification mismatch for native space " + "warping of coordinates using ANTs." + ), + ) seeds = ANTsCoordinatesWarper().warp( seeds=seeds, target_data=target_data, diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 478efba56..dddac7601 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -29,7 +29,7 @@ from ...utils import logger, raise_error from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..template_spaces import get_template -from ..utils import closest_resolution +from ..utils import closest_resolution, get_native_warper from ._ants_mask_warper import ANTsMaskWarper from ._fsl_mask_warper import FSLMaskWarper @@ -359,8 +359,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ - RuntimeError - If warper could not be found in ``extra_input``. ValueError If extra key is provided in addition to mask name in ``masks`` or if no mask is provided or @@ -374,8 +372,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): """ # Check pre-requirements for space manipulation target_space = target_data["space"] - # Set target standard space to target space - target_std_space = target_space # Extra data type requirement check if target space is native if target_space == "native": # Check for extra inputs @@ -385,16 +381,16 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) - # Set target standard space to warp file space source and warper - warper = None - for entry in extra_input["Warp"]: - if entry["dst"] == "native": - target_std_space = entry["src"] - warper = entry["warper"] - if warper is None: - raise_error( - klass=RuntimeError, msg="Could not find correct warper" - ) + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) + # Set target standard space to warp file space source + target_std_space = warper_spec["src"] + else: + # Set target standard space to target space + target_std_space = target_space # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -520,15 +516,15 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp mask if target data is native if target_space == "native": - # extra_input check done earlier and warper exists - if warper == "fsl": + # extra_input check done earlier and warper_spec exists + if warper_spec["warper"] == "fsl": mask_img = FSLMaskWarper().warp( mask_name="native", mask_img=mask_img, target_data=target_data, warp_data=warper_spec, ) - elif warper == "ants": + elif warper_spec["warper"] == "ants": mask_img = ANTsMaskWarper().warp( mask_name="native", mask_img=mask_img, diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index b8439ddb3..5362cc3cd 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -23,7 +23,7 @@ from nilearn import datasets, image from ...utils import logger, raise_error, warn_with_log from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry -from ..utils import closest_resolution +from ..utils import closest_resolution, get_native_warper from ._ants_parcellation_warper import ANTsParcellationWarper from ._fsl_parcellation_warper import FSLParcellationWarper @@ -394,16 +394,12 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ - RuntimeError - If warper could not be found in ``extra_input``. ValueError If ``extra_input`` is None when ``target_data``'s space is native. """ # Check pre-requirements for space manipulation target_space = target_data["space"] - # Set target standard space to target space - target_std_space = target_space # Extra data type requirement check if target space is native if target_space == "native": # Check for extra inputs @@ -413,16 +409,16 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) - # Set target standard space to warp file space source and warper - warper = None - for entry in extra_input["Warp"]: - if entry["dst"] == "native": - target_std_space = entry["src"] - warper = entry["warper"] - if warper is None: - raise_error( - klass=RuntimeError, msg="Could not find correct warper" - ) + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) + # Set target standard space to warp file space source + target_std_space = warper_spec["src"] + else: + # Set target standard space to target space + target_std_space = target_space # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -436,12 +432,14 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): all_parcellations = [] all_labels = [] for name in parcellations: + # Load parcellation img, labels, _, space = self.load( name=name, resolution=resolution, ) - # Convert parcellation spaces if required + # Convert parcellation spaces if required; + # cannot be "native" due to earlier check if space != target_std_space: raw_img = ANTsParcellationWarper().warp( parcellation_name=name, @@ -479,15 +477,15 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp parcellation if target space is native if target_space == "native": - # extra_input check done earlier and warper exists - if warper == "fsl": + # extra_input check done earlier and warper_spec exists + if warper_spec["warper"] == "fsl": resampled_parcellation_img = FSLParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, target_data=target_data, warp_data=warper_spec, ) - elif warper == "ants": + elif warper_spec["warper"] == "ants": resampled_parcellation_img = ANTsParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, -- 2.52.0 From e1e01cb5ca4c6b9d0b4fd8912a4d4681174a1999 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 11:56:54 +0100 Subject: [PATCH 18/20] update: correct native space condition check for warping --- junifer/data/masks/_ants_mask_warper.py | 2 +- junifer/data/masks/_masks.py | 2 +- junifer/data/parcellations/_ants_parcellation_warper.py | 2 +- junifer/data/parcellations/_parcellations.py | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index 23acd916c..90329eb57 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -83,7 +83,7 @@ class ANTsMaskWarper: ) # Native space warping - if dst == "T1w": + if dst == "native": # Warp data check if warp_data is None: raise_error("No `warp_data` provided") diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index dddac7601..8394c90ae 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -529,7 +529,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): mask_name="native", mask_img=mask_img, src="", - dst="T1w", + dst="native", target_data=target_data, warp_data=warper_spec, ) diff --git a/junifer/data/parcellations/_ants_parcellation_warper.py b/junifer/data/parcellations/_ants_parcellation_warper.py index 4d86d8c78..8eade0ce1 100644 --- a/junifer/data/parcellations/_ants_parcellation_warper.py +++ b/junifer/data/parcellations/_ants_parcellation_warper.py @@ -83,7 +83,7 @@ class ANTsParcellationWarper: ) # Native space warping - if dst == "T1w": + if dst == "native": # Warp data check if warp_data is None: raise_error("No `warp_data` provided") diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 5362cc3cd..129b13226 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -490,7 +490,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): parcellation_name="native", parcellation_img=resampled_parcellation_img, src="", - dst="T1w", + dst="native", target_data=target_data, warp_data=warper_spec, ) -- 2.52.0 From 6125bade46958510a6b3c49b5f7296e851b6c1e7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 12:19:31 +0100 Subject: [PATCH 19/20] fix: use "brain" as template instead of "T1w" for resampling in compute_brain_mask() --- junifer/data/masks/_masks.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 8394c90ae..a868ccf30 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -114,7 +114,7 @@ def compute_brain_mask( space=target_std_space, target_data=target_data, extra_input=extra_input, - template_type=mask_type if mask_type in ["gm", "wm"] else "T1w", + template_type=mask_type, ) # Resample template to target image target_img = target_data["data"] -- 2.52.0 From f5e5792a24302f7055823ea82c0bbf5540e2fe43 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 13 Nov 2024 12:27:07 +0100 Subject: [PATCH 20/20] chore: add changelogs 390.{change,enh} --- docs/changes/newsfragments/390.change | 1 + docs/changes/newsfragments/390.enh | 1 + 2 files changed, 2 insertions(+) create mode 100644 docs/changes/newsfragments/390.change create mode 100644 docs/changes/newsfragments/390.enh diff --git a/docs/changes/newsfragments/390.change b/docs/changes/newsfragments/390.change new file mode 100644 index 000000000..6356751c2 --- /dev/null +++ b/docs/changes/newsfragments/390.change @@ -0,0 +1 @@ +:class:`.SpaceWarper`'s ``using`` now supports ``"auto"`` allowing auto tool selection based on DataGrabber specification by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/390.enh b/docs/changes/newsfragments/390.enh new file mode 100644 index 000000000..96545d4b6 --- /dev/null +++ b/docs/changes/newsfragments/390.enh @@ -0,0 +1 @@ +Update ``Warp`` data type to be a list of dictionaries and adapt to allow multiple transforms to be specified by `Synchon Mandal`_ -- 2.52.0