diff --git a/docs/builtin.rst b/docs/builtin.rst index 2f12e04ae..bc6876bef 100644 --- a/docs/builtin.rst +++ b/docs/builtin.rst @@ -192,6 +192,14 @@ Available as found in `Jo et al. (2021) `_ - Done - 0.0.2 + * - :class:`junifer.markers.TemporalSNRParcels` + - Calculate temporal signal-to-noise ratio using parcellations + - Done + - 0.0.2 + * - :class:`junifer.markers.TemporalSNRSpheres` + - Calculate temporal signal-to-noise ratio using spheres placed on coordinates + - Done + - 0.0.2 Planned ~~~~~~~ diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index a54440683..ff0da3932 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -43,6 +43,9 @@ Enhancements - Add fMRIPrep brain masks to the datagrabber patterns for all datagrabbers in the aomic sub-package (:gh:`177` by `Leonard Sasse`_). +- Add :class:`junifer.markers.TemporalSNRParcels` and :class:`junifer.markers.TemporalSNRSpheres` + (:gh:`163` by `Leonard Sasse`_). + Bugs ~~~~ diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index badbd2f30..26f259edf 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -22,3 +22,7 @@ from .falff import ( AmplitudeLowFrequencyFluctuationParcels, AmplitudeLowFrequencyFluctuationSpheres, ) +from .temporal_snr import ( + TemporalSNRParcels, + TemporalSNRSpheres, +) diff --git a/junifer/markers/temporal_snr/__init__.py b/junifer/markers/temporal_snr/__init__.py new file mode 100644 index 000000000..ca150f982 --- /dev/null +++ b/junifer/markers/temporal_snr/__init__.py @@ -0,0 +1,7 @@ +"""Provide imports for temporal signal-to-noise ratio sub-package.""" + +# Authors: Leonard Sasse +# License: AGPL + +from .temporal_snr_parcels import TemporalSNRParcels +from .temporal_snr_spheres import TemporalSNRSpheres diff --git a/junifer/markers/temporal_snr/temporal_snr_base.py b/junifer/markers/temporal_snr/temporal_snr_base.py new file mode 100644 index 000000000..7679c6c8e --- /dev/null +++ b/junifer/markers/temporal_snr/temporal_snr_base.py @@ -0,0 +1,128 @@ +"""Provide abstract base class for temporal signal-to-noise ratio (tSNR).""" + +# Authors: Leonard Sasse +# License: AGPL + + +from abc import abstractmethod +from typing import Any, Dict, List, Optional, Union + +from nilearn import image as nimg + +from ...utils import raise_error +from ..base import BaseMarker + + +class TemporalSNRBase(BaseMarker): + """Abstract base class for temporal SNR markers. + + Parameters + ---------- + agg_method : str, optional + The method to perform aggregation using. Check valid options in + :func:`junifer.stats.get_aggfunc_by_name` (default "mean"). + agg_method_params : dict, optional + Parameters to pass to the aggregation function. Check valid options in + :func:`junifer.stats.get_aggfunc_by_name` (default None). + masks : str, dict or list of dict or str, optional + The specification of the masks to apply to regions before extracting + signals. Check :ref:`Using Masks ` for more details. + If None, will not apply any mask (default None). + name : str, optional + The name of the marker. If None, will use the class name (default + None). + + """ + + _DEPENDENCIES = {"nilearn"} + + def __init__( + self, + agg_method: str = "mean", + agg_method_params: Optional[Dict] = None, + masks: Union[str, Dict, List[Union[Dict, str]], None] = None, + name: Optional[str] = None, + ) -> None: + self.agg_method = agg_method + self.agg_method_params = agg_method_params + self.masks = masks + super().__init__(on="BOLD", name=name) + + @abstractmethod + def aggregate( + self, input: Dict[str, Any], extra_input: Optional[Dict] = None + ) -> Dict[str, Any]: + """Perform aggregation.""" + raise_error( + msg="Concrete classes need to implement aggregate().", + klass=NotImplementedError, + ) + + def get_valid_inputs(self) -> List[str]: + """Get valid data types for input. + + Returns + ------- + list of str + The list of data types that can be used as input for this marker. + """ + return ["BOLD"] + + def get_output_type(self, input_type: str) -> str: + """Get output type. + + Parameters + ---------- + input_type : str + The data type input to the marker. + + Returns + ------- + str + The storage type output by the marker. + + """ + return "vector" + + def compute( + self, + input: Dict[str, Any], + extra_input: Optional[Dict] = None, + ) -> Dict: + """Compute. + + 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. + + Returns + ------- + dict + The computed result as dictionary. The following keys will be + included in the dictionary: + + * ``data`` : the computed values as a ``numpy.ndarray`` + * ``col_names`` : the column labels for the computed values as list + + """ + # Calculate voxelwise temporal signal-to-noise ratio in an image + mean_img = nimg.math_img( + "img.mean(axis=-1).squeeze()", img=input["data"] + ) + stdv_img = nimg.math_img( + "img.std(axis=-1).squeeze()", img=input["data"] + ) + mask_img = nimg.math_img("(stdv_img != 0)", stdv_img=stdv_img) + input["data"] = nimg.math_img( + "np.divide(mean_img, stdv_img, where=mask_img.astype(bool))", + mean_img=mean_img, + stdv_img=stdv_img, + mask_img=mask_img, + ) + # Perform necessary aggregation and return + return self.aggregate(input=input, extra_input=extra_input) diff --git a/junifer/markers/temporal_snr/temporal_snr_parcels.py b/junifer/markers/temporal_snr/temporal_snr_parcels.py new file mode 100644 index 000000000..0521527de --- /dev/null +++ b/junifer/markers/temporal_snr/temporal_snr_parcels.py @@ -0,0 +1,89 @@ +"""Provide class for temporal SNR using parcels.""" + +# Authors: Leonard Sasse +# License: AGPL + +from typing import Any, Dict, List, Optional, Union + +from ...api.decorators import register_marker +from ..parcel_aggregation import ParcelAggregation +from .temporal_snr_base import TemporalSNRBase + + +@register_marker +class TemporalSNRParcels(TemporalSNRBase): + """Class for temporal signal-to-noise ratio using parcellations. + + Parameters + ---------- + parcellation : str or list of str + The name(s) of the parcellation(s). Check valid options by calling + :func:`junifer.data.parcellations.list_parcellations`. + agg_method : str, optional + The method to perform aggregation using. Check valid options in + :func:`junifer.stats.get_aggfunc_by_name` (default "mean"). + agg_method_params : dict, optional + Parameters to pass to the aggregation function. Check valid options in + :func:`junifer.stats.get_aggfunc_by_name` (default None). + masks : str, dict or list of dict or str, optional + The specification of the masks to apply to regions before extracting + signals. Check :ref:`Using Masks ` for more details. + If None, will not apply any mask (default None). + name : str, optional + The name of the marker. If None, will use the class name (default + None). + + """ + + def __init__( + self, + parcellation: Union[str, List[str]], + agg_method: str = "mean", + agg_method_params: Optional[Dict] = None, + masks: Union[str, Dict, List[Union[Dict, str]], None] = None, + name: Optional[str] = None, + ) -> None: + self.parcellation = parcellation + super().__init__( + agg_method=agg_method, + agg_method_params=agg_method_params, + masks=masks, + name=name, + ) + + 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 the data + is the voxelwise temporal SNR map. + 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`` : ROI-wise temporal SNR as a ``numpy.ndarray`` + * ``col_names`` : the ROI labels for the computed values as list + + """ + parcel_aggregation = ParcelAggregation( + parcellation=self.parcellation, + method=self.agg_method, + method_params=self.agg_method_params, + masks=self.masks, + on="BOLD", + ) + # Return the 2D timeseries after parcel aggregation + return parcel_aggregation.compute(input=input, extra_input=extra_input) diff --git a/junifer/markers/temporal_snr/temporal_snr_spheres.py b/junifer/markers/temporal_snr/temporal_snr_spheres.py new file mode 100644 index 000000000..402d4aa47 --- /dev/null +++ b/junifer/markers/temporal_snr/temporal_snr_spheres.py @@ -0,0 +1,100 @@ +"""Provide class for temporal SNR using spheres.""" + +# Authors: Leonard Sasse +# License: AGPL + +from typing import Any, Dict, List, Optional, Union + +from ...api.decorators import register_marker +from ..sphere_aggregation import SphereAggregation +from ..utils import raise_error +from .temporal_snr_base import TemporalSNRBase + + +@register_marker +class TemporalSNRSpheres(TemporalSNRBase): + """Class for temporal signal-to-noise ratio using coordinates (spheres). + + Parameters + ---------- + coords : str + The name of the coordinates list to use. See + :func:`junifer.data.coordinates.list_coordinates` for options. + radius : float, optional + The radius of the sphere in mm. If None, the signal will be extracted + from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` + for more information (default None). + agg_method : str, optional + The aggregation method to use. + See :func:`junifer.stats.get_aggfunc_by_name` for more information + (default None). + agg_method_params : dict, optional + The parameters to pass to the aggregation method (default None). + masks : str, dict or list of dict or str, optional + The specification of the masks to apply to regions before extracting + signals. Check :ref:`Using Masks ` for more details. + If None, will not apply any mask (default None). + name : str, optional + The name of the marker. By default, it will use + KIND_FunctionalConnectivitySpheres where KIND is the kind of data it + was applied to (default None). + + """ + + def __init__( + self, + coords: str, + radius: Optional[float] = None, + agg_method: str = "mean", + agg_method_params: Optional[Dict] = None, + masks: Union[str, Dict, List[Union[Dict, str]], None] = None, + name: Optional[str] = None, + ) -> None: + self.coords = coords + self.radius = radius + if radius is None or radius <= 0: + raise_error(f"radius should be > 0: provided {radius}") + super().__init__( + agg_method=agg_method, + agg_method_params=agg_method_params, + masks=masks, + name=name, + ) + + 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 the data + is the voxelwise temporal SNR map. + 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`` : VOI-wise temporal SNR as a ``numpy.ndarray`` + * ``col_names`` : the VOI labels for the computed values as list + + """ + sphere_aggregation = SphereAggregation( + coords=self.coords, + radius=self.radius, + method=self.agg_method, + method_params=self.agg_method_params, + masks=self.masks, + on="BOLD", + ) + # Return the 2D timeseries after sphere aggregation + return sphere_aggregation.compute(input=input, extra_input=extra_input) diff --git a/junifer/markers/temporal_snr/tests/test_temporal_snr_base.py b/junifer/markers/temporal_snr/tests/test_temporal_snr_base.py new file mode 100644 index 000000000..396dd6300 --- /dev/null +++ b/junifer/markers/temporal_snr/tests/test_temporal_snr_base.py @@ -0,0 +1,15 @@ +"""Provide tests for base temporal SNR marker.""" + +# Authors: Leonard Sasse +# License: AGPL + +import pytest + +# done to keep line length 79 +import junifer.markers.temporal_snr as tsnr + + +def test_base_temporal_snr_marker_abstractness() -> None: + """Test TemporalSNRBase is an abstract base class.""" + with pytest.raises(TypeError, match="abstract"): + tsnr.temporal_snr_base.TemporalSNRBase() diff --git a/junifer/markers/temporal_snr/tests/test_temporal_snr_parcels.py b/junifer/markers/temporal_snr/tests/test_temporal_snr_parcels.py new file mode 100644 index 000000000..e0ef668f4 --- /dev/null +++ b/junifer/markers/temporal_snr/tests/test_temporal_snr_parcels.py @@ -0,0 +1,54 @@ +"""Provide tests for temporal signal-to-noise ratio using parcellation.""" + +# Authors: Leonard Sasse +# License: AGPL + +from pathlib import Path + +from nilearn import datasets, image + +from junifer.markers.temporal_snr import TemporalSNRParcels +from junifer.storage import SQLiteFeatureStorage + + +def test_TemporalSNRParcels(tmp_path: Path) -> None: + """Test TemporalSNRParcels. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # get a dataset + ni_data = datasets.fetch_spm_auditory(subject_id="sub001") + fmri_img = image.concat_imgs(ni_data.func) # type: ignore + + tsnr_parcels = TemporalSNRParcels(parcellation="Schaefer100x7") + all_out = tsnr_parcels.fit_transform( + {"BOLD": {"data": fmri_img, "meta": {}}} + ) + + out = all_out["BOLD"] + + assert "data" in out + assert "col_names" in out + + assert out["data"].shape[0] == 1 + assert out["data"].shape[1] == 100 + assert len(set(out["col_names"])) == 100 + + # check correct output + assert tsnr_parcels.get_output_type("BOLD") == "vector" + + uri = tmp_path / "test_tsnr_parcellation.sqlite" + # Single storage, must be the uri + storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") + meta = {"element": {"subject": "test"}, "dependencies": {"numpy"}} + input = {"BOLD": {"data": fmri_img, "meta": meta}} + all_out = tsnr_parcels.fit_transform(input, storage=storage) + + features = storage.list_features() + assert any( + x["name"] == "BOLD_TemporalSNRParcels" for x in features.values() + ) diff --git a/junifer/markers/temporal_snr/tests/test_temporal_snr_spheres.py b/junifer/markers/temporal_snr/tests/test_temporal_snr_spheres.py new file mode 100644 index 000000000..448de5e54 --- /dev/null +++ b/junifer/markers/temporal_snr/tests/test_temporal_snr_spheres.py @@ -0,0 +1,63 @@ +"""Provide tests for temporal signal-to-noise spheres.""" + +# Authors: Leonard Sasse +# License: AGPL + +from pathlib import Path + +import pytest +from nilearn import datasets, image + +from junifer.markers.temporal_snr import TemporalSNRSpheres +from junifer.storage import SQLiteFeatureStorage + + +def test_TemporalSNRSpheres(tmp_path: Path) -> None: + """Test TemporalSNRSpheres. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # get a dataset + ni_data = datasets.fetch_spm_auditory(subject_id="sub001") + fmri_img = image.concat_imgs(ni_data.func) # type: ignore + + tsnr_spheres = TemporalSNRSpheres(coords="DMNBuckner", radius=5.0) + all_out = tsnr_spheres.fit_transform( + {"BOLD": {"data": fmri_img, "meta": {}}} + ) + + out = all_out["BOLD"] + + assert "data" in out + assert "col_names" in out + assert out["data"].shape[0] == 1 + assert out["data"].shape[1] == 6 + assert len(set(out["col_names"])) == 6 + + # check correct output + assert tsnr_spheres.get_output_type("BOLD") == "vector" + + uri = tmp_path / "test_tsnr_coords.sqlite" + # Single storage, must be the uri + storage = SQLiteFeatureStorage(uri=uri, upsert="ignore") + meta = { + "element": {"subject": "test"}, + "dependencies": {"numpy", "nilearn"}, + } + input = {"BOLD": {"data": fmri_img, "meta": meta}} + all_out = tsnr_spheres.fit_transform(input, storage=storage) + + features = storage.list_features() + assert any( + x["name"] == "BOLD_TemporalSNRSpheres" for x in features.values() + ) + + +def test_TemporalSNRSpheres_error() -> None: + """Test TemporalSNRSpheres errors.""" + with pytest.raises(ValueError, match="radius should be > 0"): + TemporalSNRSpheres(coords="DMNBuckner", radius=-0.1)