Add missing extra_input parameters #189

Merged
fraimondo merged 2 commits from fix/extra_input_2 into main 2023-03-15 15:04:02 +00:00
6 changed files with 119 additions and 14 deletions

View file

@ -57,6 +57,8 @@ Bugs
- 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`_).
- Fix several markers that did not properly handle the ``extra_input`` parameter (:gh:`189` by `Fede Raimondo`_).
API changes
~~~~~~~~~~~

View file

@ -130,12 +130,12 @@ class CrossParcellationFC(BaseMarker):
parcellation=self.parcellation_one,
method=self.aggregation_method,
masks=self.masks,
).compute(input)
).compute(input, extra_input=extra_input)
parcellation_two_dict = ParcelAggregation(
parcellation=self.parcellation_two,
method=self.aggregation_method,
masks=self.masks,
).compute(input)
).compute(input, extra_input=extra_input)
parcellated_ts_one = parcellation_one_dict["data"]
parcellated_ts_two = parcellation_two_dict["data"]

View file

@ -72,8 +72,32 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform parcel aggregation."""
def aggregate(
self, input: Dict[str, Any], extra_input: Optional[Dict] = None
) -> Dict:
"""Perform parcel aggregation and ETS computation.
Parameters
----------
input : dict
A single input from the pipeline data object in which to compute
the marker.
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).
Returns
-------
dict
The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``col_names`` : the column labels for the computed values as list
"""
parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
@ -82,7 +106,9 @@ class EdgeCentricFCParcels(FunctionalConnectivityBase):
on="BOLD",
)
bold_aggregated = parcel_aggregation.compute(input)
bold_aggregated = parcel_aggregation.compute(
input, extra_input=extra_input
)
ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["col_names"]
)

View file

@ -79,8 +79,33 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform sphere aggregation."""
def aggregate(
self, input: Dict[str, Any], extra_input: Optional[Dict] = None
) -> Dict:
"""Perform sphere aggregation and ETS computation.
Parameters
----------
input : dict
A single input from the pipeline data object in which to compute
the marker.
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).
Returns
-------
dict
The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``col_names`` : the column labels for the computed values as list
"""
sphere_aggregation = SphereAggregation(
coords=self.coords,
radius=self.radius,
@ -89,7 +114,9 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
masks=self.masks,
on="BOLD",
)
bold_aggregated = sphere_aggregation.compute(input)
bold_aggregated = sphere_aggregation.compute(
input, extra_input=extra_input
)
ets, edge_names = _ets(
bold_aggregated["data"], bold_aggregated["col_names"]
)

View file

@ -64,8 +64,33 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform parcel aggregation."""
def aggregate(
self, input: Dict[str, Any], extra_input: Optional[Dict] = None
) -> Dict:
"""Perform parcel aggregation.
Parameters
----------
input : dict
A single input from the pipeline data object in which to compute
the marker.
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).
Returns
-------
dict
The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``col_names`` : the column labels for the computed values as list
"""
parcel_aggregation = ParcelAggregation(
parcellation=self.parcellation,
method=self.agg_method,
@ -74,4 +99,4 @@ class FunctionalConnectivityParcels(FunctionalConnectivityBase):
on="BOLD",
)
# Return the 2D timeseries after parcel aggregation
return parcel_aggregation.compute(input)
return parcel_aggregation.compute(input, extra_input=extra_input)

View file

@ -73,8 +73,33 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
name=name,
)
def aggregate(self, input: Dict[str, Any]) -> Dict:
"""Perform sphere aggregation."""
def aggregate(
self, input: Dict[str, Any], extra_input: Optional[Dict] = None
) -> Dict:
"""Perform sphere aggregation.
Parameters
----------
input : dict
A single input from the pipeline data object in which to compute
the marker.
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).
Returns
-------
dict
The computed result as dictionary. This will be either returned
to the user or stored in the storage by calling the store method
with this as a parameter. The dictionary has the following keys:
* ``data`` : the actual computed values as a numpy.ndarray
* ``col_names`` : the column labels for the computed values as list
"""
sphere_aggregation = SphereAggregation(
coords=self.coords,
radius=self.radius,
@ -84,4 +109,4 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
on="BOLD",
)
# Return the 2D timeseries after sphere aggregation
return sphere_aggregation.compute(input)
return sphere_aggregation.compute(input, extra_input=extra_input)