Add extra_input to reho/falff #187
6 changed files with 34 additions and 13 deletions
|
|
@ -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
|
||||
~~~~~~~~~~~
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue