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 ``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
|
||||||
~~~~~~~~~~~
|
~~~~~~~~~~~
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue