Add extra_input to reho/falff #187

Merged
fraimondo merged 2 commits from fix/extra_input into main 2023-03-14 22:11:37 +00:00
6 changed files with 34 additions and 13 deletions

View file

@ -54,6 +54,9 @@ Bugs
- Fix ``junifer run`` to respect preprocess step specified in the pipeline (:gh:`159` by `Synchon Mandal`_). - Fix ``junifer run`` to respect preprocess step specified in the pipeline (:gh:`159` by `Synchon Mandal`_).
- Fix :class:`junifer.markers.AmplitudeLowFrequencyFluctuationParcels`, :class:`junifer.markers.AmplitudeLowFrequencyFluctuationSpheres`
:class:`junifer.markers.ReHoSpheres` and :class:`junifer.markers.ReHoParcels` pass the ``extra_input`` parameter (:gh:`187` by `Fede Raimondo`_).
API changes API changes
~~~~~~~~~~~ ~~~~~~~~~~~

View file

@ -61,8 +61,8 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
use_afni: Optional[bool] = None, use_afni: Optional[bool] = None,
name: Optional[str] = None, name: Optional[str] = None,
) -> None: ) -> None:
if highpass <= 0: if highpass < 0:
raise_error("Highpass must be positive") raise_error("Highpass must be positive or 0")
if lowpass <= 0: if lowpass <= 0:
raise_error("Lowpass must be positive") raise_error("Lowpass must be positive")
if highpass >= lowpass: if highpass >= lowpass:
@ -160,18 +160,25 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
"path": None, "path": None,
} }
out = self._postprocess(post_input) out = self._postprocess(post_input, extra_input=extra_input)
return out return out
@abstractmethod @abstractmethod
def _postprocess(self, input: Dict) -> Dict: def _postprocess(
self, input: Dict, extra_input: Optional[Dict] = None
) -> Dict:
"""Postprocess the output of the estimator. """Postprocess the output of the estimator.
Parameters Parameters
---------- ----------
input : dict input : dict
The output of the estimator. It must have the following The output of the estimator. It must have the following
extra_input : dict, optional
The other fields in the pipeline data object. Useful for accessing
other data kind that needs to be used in the computation. For
example, the functional connectivity markers can make use of the
confounds if available (default None).
""" """
raise_error( raise_error(
"_postprocess must be implemented", klass=NotImplementedError "_postprocess must be implemented", klass=NotImplementedError

View file

@ -26,7 +26,8 @@ class AmplitudeLowFrequencyFluctuationParcels(
fractional : bool fractional : bool
Whether to compute fractional ALFF. Whether to compute fractional ALFF.
highpass : positive float, optional highpass : positive float, optional
The highpass cutoff frequency for the bandpass filter (default 0.01). The highpass cutoff frequency for the bandpass filter. If 0,
it will not apply a highpass filter (default 0.01).
lowpass : positive float, optional lowpass : positive float, optional
The lowpass cutoff frequency for the bandpass filter (default 0.1). The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional tr : positive float, optional
@ -88,7 +89,9 @@ class AmplitudeLowFrequencyFluctuationParcels(
use_afni=use_afni, use_afni=use_afni,
) )
def _postprocess(self, input: Dict) -> Dict: def _postprocess(
self, input: Dict, extra_input: Optional[Dict] = None
) -> Dict:
"""Compute ALFF and fALFF. """Compute ALFF and fALFF.
Parameters Parameters
@ -121,6 +124,6 @@ class AmplitudeLowFrequencyFluctuationParcels(
) )
# get the 2D timeseries after parcel aggregation # get the 2D timeseries after parcel aggregation
out = pa.compute(input) out = pa.compute(input, extra_input=extra_input)
return out return out

View file

@ -30,7 +30,8 @@ class AmplitudeLowFrequencyFluctuationSpheres(
fractional : bool fractional : bool
Whether to compute fractional ALFF. Whether to compute fractional ALFF.
highpass : positive float, optional highpass : positive float, optional
The highpass cutoff frequency for the bandpass filter (default 0.01). The highpass cutoff frequency for the bandpass filter. If 0,
it will not apply a highpass filter (default 0.01).
lowpass : positive float, optional lowpass : positive float, optional
The lowpass cutoff frequency for the bandpass filter (default 0.1). The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional tr : positive float, optional
@ -94,7 +95,9 @@ class AmplitudeLowFrequencyFluctuationSpheres(
use_afni=use_afni, use_afni=use_afni,
) )
def _postprocess(self, input: Dict) -> Dict: def _postprocess(
self, input: Dict, extra_input: Optional[Dict] = None
) -> Dict:
"""Compute ALFF and fALFF. """Compute ALFF and fALFF.
Parameters Parameters
@ -127,7 +130,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
on="fALFF", on="fALFF",
) )
# get the 2D timeseries after parcel aggregation # get the 2D timeseries after sphere aggregation
out = pa.compute(input) out = pa.compute(input, extra_input=extra_input)
return out return out

View file

@ -141,7 +141,10 @@ class ReHoParcels(ReHoBase):
) )
# Perform aggregation on reho map # Perform aggregation on reho map
parcel_aggregation_input = {"data": reho_map} parcel_aggregation_input = {"data": reho_map}
output = parcel_aggregation.compute(input=parcel_aggregation_input) output = parcel_aggregation.compute(
input=parcel_aggregation_input,
extra_input=extra_input,
)
# Only use the first row and expand row dimension # Only use the first row and expand row dimension
output["data"] = output["data"][0][np.newaxis, :] output["data"] = output["data"][0][np.newaxis, :]
return output return output

View file

@ -149,7 +149,9 @@ class ReHoSpheres(ReHoBase):
) )
# Perform aggregation on reho map # Perform aggregation on reho map
sphere_aggregation_input = {"data": reho_map} sphere_aggregation_input = {"data": reho_map}
output = sphere_aggregation.compute(input=sphere_aggregation_input) output = sphere_aggregation.compute(
input=sphere_aggregation_input, extra_input=extra_input
)
# Only use the first row and expand row dimension # Only use the first row and expand row dimension
output["data"] = output["data"][0][np.newaxis, :] output["data"] = output["data"][0][np.newaxis, :]
return output return output