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 :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
~~~~~~~~~~~

View file

@ -61,8 +61,8 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
use_afni: Optional[bool] = None,
name: Optional[str] = None,
) -> None:
if highpass <= 0:
raise_error("Highpass must be positive")
if highpass < 0:
raise_error("Highpass must be positive or 0")
if lowpass <= 0:
raise_error("Lowpass must be positive")
if highpass >= lowpass:
@ -160,18 +160,25 @@ class AmplitudeLowFrequencyFluctuationBase(BaseMarker):
"path": None,
}
out = self._postprocess(post_input)
out = self._postprocess(post_input, extra_input=extra_input)
return out
@abstractmethod
def _postprocess(self, input: Dict) -> Dict:
def _postprocess(
self, input: Dict, extra_input: Optional[Dict] = None
) -> Dict:
"""Postprocess the output of the estimator.
Parameters
----------
input : dict
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(
"_postprocess must be implemented", klass=NotImplementedError

View file

@ -26,7 +26,8 @@ class AmplitudeLowFrequencyFluctuationParcels(
fractional : bool
Whether to compute fractional ALFF.
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
The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional
@ -88,7 +89,9 @@ class AmplitudeLowFrequencyFluctuationParcels(
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.
Parameters
@ -121,6 +124,6 @@ class AmplitudeLowFrequencyFluctuationParcels(
)
# get the 2D timeseries after parcel aggregation
out = pa.compute(input)
out = pa.compute(input, extra_input=extra_input)
return out

View file

@ -30,7 +30,8 @@ class AmplitudeLowFrequencyFluctuationSpheres(
fractional : bool
Whether to compute fractional ALFF.
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
The lowpass cutoff frequency for the bandpass filter (default 0.1).
tr : positive float, optional
@ -94,7 +95,9 @@ class AmplitudeLowFrequencyFluctuationSpheres(
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.
Parameters
@ -127,7 +130,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
on="fALFF",
)
# get the 2D timeseries after parcel aggregation
out = pa.compute(input)
# get the 2D timeseries after sphere aggregation
out = pa.compute(input, extra_input=extra_input)
return out

View file

@ -141,7 +141,10 @@ class ReHoParcels(ReHoBase):
)
# Perform aggregation on 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
output["data"] = output["data"][0][np.newaxis, :]
return output

View file

@ -149,7 +149,9 @@ class ReHoSpheres(ReHoBase):
)
# Perform aggregation on 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
output["data"] = output["data"][0][np.newaxis, :]
return output