diff --git a/.codespellrc b/.codespellrc deleted file mode 100644 index ddfbbf2f8..000000000 --- a/.codespellrc +++ /dev/null @@ -1,7 +0,0 @@ -[codespell] -skip = docs/auto_*,*.html,.git/,*.pyc,docs/_build -count = -quiet-level = 3 -ignore-words = ignore_words.txt -interactive = 0 -builtin = clear,rare,informal,names,usage \ No newline at end of file diff --git a/.coveragerc b/.coveragerc deleted file mode 100644 index f7d2883cb..000000000 --- a/.coveragerc +++ /dev/null @@ -1,12 +0,0 @@ -[run] -branch = True -source = junifer -include = */junifer/* -omit = - */setup.py - */tests/* - -[report] -exclude_lines = - pragma: no cover - if __name__ == .__main__.: \ No newline at end of file diff --git a/.flake8 b/.flake8 deleted file mode 100644 index d0ed0eec7..000000000 --- a/.flake8 +++ /dev/null @@ -1,5 +0,0 @@ -[flake8] -exclude = __init__.py,*externals*,constants.py,fixes.py,resources.py,nilearn_cache,venv,docs/auto_examples,docs/_build/,.eggs/ -ignore = W503,W504,I100,I101,I201,N806,E201,E202,E221,E222,E241,F541 -# We add A for the array-spacing plugin, and ignore the E ones it covers above -select = A,E,F,W,C \ No newline at end of file diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 344ac188a..b7cbe6d0d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,52 +4,36 @@ on: [push, pull_request] jobs: build: - runs-on: ubuntu-latest strategy: + fail-fast: false matrix: - python-version: [3.6, 3.7, 3.8] + python-version: ['3.8', '3.9', '3.10'] steps: - - uses: actions/checkout@v2 - - name: Set up Python ${{ matrix.python-version }} - uses: actions/setup-python@v2 - with: - python-version: ${{ matrix.python-version }} - - name: Check for sudo - shell: bash + - name: Set up system run: | - if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi - echo "SUDO=$SUDO" >> $GITHUB_ENV - - name: Install dependencies - run: | - $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" - $SUDO apt-get update -qq - $SUDO apt-get install git-annex-standalone - python -m pip install --upgrade pip - pip install -r test-requirements.txt - pip install -r requirements.txt + bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" + sudo apt-get update -qq + sudo apt-get install git-annex-standalone - name: Configure git for datalad run: | git config --global user.email "runner@github.com" - git config --global user.name "GITHUB CI Runner" - - name: Install junifer - shell: bash -el {0} + git config --global user.name "GitHub Runner" + - uses: actions/checkout@v3 + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v4 + with: + python-version: ${{ matrix.python-version }} + - name: Install dependencies run: | - python setup.py build - python setup.py install - - name: Lint with flake8 + python -m pip install --upgrade pip setuptools wheel + python -m pip install tox tox-gh-actions + - name: Test with tox run: | - # stop the build if there are Python syntax errors or undefined names - flake8 . --count --show-source --statistics - - name: Spell check - run: | - codespell junifer/ docs/ examples/ - - name: Test with pytest - run: | - PYTHONPATH="." pytest --cov=junifer --cov-report xml -vv junifer/ - - name: 'Upload coverage to CodeCov' - uses: codecov/codecov-action@master + tox + - name: Upload coverage to Codecov + uses: codecov/codecov-action@v3 with: token: ${{ secrets.CODECOV_TOKEN }} - if: success() && matrix.python-version == 3.8 + if: success() && matrix.python-version == 3.9 diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index 58b7314ef..64a033b2c 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -7,64 +7,58 @@ jobs: runs-on: ubuntu-latest strategy: fail-fast: false - steps: - - name: Checkout Source - uses: actions/checkout@v2 - with: - # require all of history to see all tagged versions' docs - fetch-depth: 0 - - name: Set up Python ${{ matrix.python-version }} - uses: actions/setup-python@v2 - with: - python-version: 3.8 - - name: Check for sudo - shell: bash - run: | - if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi - echo "SUDO=$SUDO" >> $GITHUB_ENV - - name: Install Dependencies - run: | - $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" - $SUDO apt-get update -qq - $SUDO apt-get install git-annex-standalone - python -m pip install --upgrade pip - pip install -r requirements.txt - pip install -r docs-requirements.txt - python setup.py build - python setup.py install - - name: Configure git for datalad - run: | - git config --global user.email "runner@github.com" - git config --global user.name "GITHUB CI Runner" - - name: Checkout gh-pages - # As we already did a deploy of gh-pages above, it is guaranteed to be there - # so check it out so we can selectively build docs below - uses: actions/checkout@v2 - with: + steps: + - name: Checkout source + uses: actions/checkout@v2 + with: + # require all of history to see all tagged versions' docs + fetch-depth: 0 + - name: Set up Python 3.9 + uses: actions/setup-python@v2 + with: + python-version: 3.9 + - name: Check for sudo + shell: bash + run: | + if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi + echo "SUDO=$SUDO" >> $GITHUB_ENV + - name: Install dependencies + run: | + $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" + $SUDO apt-get update -qq + $SUDO apt-get install git-annex-standalone + python -m pip install --upgrade pip setuptools wheel + python -m pip install -e .[docs] + - name: Configure git for datalad + run: | + git config --global user.email "runner@github.com" + git config --global user.name "GITHUB CI Runner" + - name: Checkout gh-pages + # As we already did a deploy of gh-pages above, it is guaranteed to be there + # so check it out so we can selectively build docs below + uses: actions/checkout@v2 + with: ref: gh-pages path: docs/_build - - - name: Test Build Docs - if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags') - run: | - BUILDDIR=_build/main make -C docs/ local - - - name: Build Docs - if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') - # Use the args we normally pass to sphinx-build, but run sphinx-multiversion - run: | - make -C docs/ html - touch docs/_build/.nojekyll - cp docs/redirect.html docs/_build/index.html - - - name: Publish Docs to gh-pages - # Only once from main or a tag - if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') - # We pin to the SHA, not the tag, for security reasons. - # https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions - uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3 - with: - github_token: ${{ secrets.GITHUB_TOKEN }} - publish_dir: docs/_build - keep_files: true + - name: Test build docs + if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags') + run: | + BUILDDIR=_build/main make -C docs/ local + - name: Build docs + if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') + # Use the args we normally pass to sphinx-build, but run sphinx-multiversion + run: | + make -C docs/ html + touch docs/_build/.nojekyll + cp docs/redirect.html docs/_build/index.html + - name: Publish docs to gh-pages + # Only once from main or a tag + if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') + # We pin to the SHA, not the tag, for security reasons. + # https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions + uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3 + with: + github_token: ${{ secrets.GITHUB_TOKEN }} + publish_dir: docs/_build + keep_files: true diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml new file mode 100644 index 000000000..edf01ba2e --- /dev/null +++ b/.github/workflows/lint.yml @@ -0,0 +1,31 @@ +name: Lint + +on: + - push + - pull_request + +jobs: + lint: + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest] + python-version: ['3.10'] + + steps: + - uses: actions/checkout@v3 + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v4 + with: + python-version: ${{ matrix.python-version }} + - name: Install dependencies + run: | + python -m pip install --upgrade pip setuptools wheel + python -m pip install tox tox-gh-actions + - name: Check with flake8 + run: | + tox -e flake8 + - name: Check with codespell + run: | + tox -e codespell diff --git a/.gitignore b/.gitignore index 09d12056a..94ac529c9 100644 --- a/.gitignore +++ b/.gitignore @@ -132,4 +132,6 @@ cython_debug/ # OS Stuff .DS_store -junifer/_version.py \ No newline at end of file +junifer/_version.py +scratch/ +junifer_jobs/ \ No newline at end of file diff --git a/AUTHORS.rst b/AUTHORS.rst index ac9115336..36fbf49de 100644 --- a/AUTHORS.rst +++ b/AUTHORS.rst @@ -1,4 +1,5 @@ Original Authors ================ * Federico Raimondo -* Leonard Sasse \ No newline at end of file +* Leonard Sasse +* Synchon Mandal diff --git a/Makefile b/Makefile deleted file mode 100644 index f34d0536d..000000000 --- a/Makefile +++ /dev/null @@ -1,15 +0,0 @@ -# Makefile before PR -# - -.PHONY: checks - -checks: flake spellcheck - -flake: - flake8 - -spellcheck: - codespell junifer/ docs/ examples/ - -test: - pytest -v \ No newline at end of file diff --git a/README.md b/README.md index 64d11f49a..31ed94d14 100644 --- a/README.md +++ b/README.md @@ -1,20 +1,68 @@ -# python-library-mockup -JUelich NeuroImaging FEature extractoR +# junifer - JUelich NeuroImaging FEature extractoR + +![PyPI](https://img.shields.io/pypi/v/junifer?style=flat-square) +![PyPI - Python Version](https://img.shields.io/pypi/pyversions/junifer?style=flat-square) +![PyPI - Wheel](https://img.shields.io/pypi/wheel/junifer?style=flat-square) +![GitHub](https://img.shields.io/github/license/juaml/junifer?style=flat-square) +[![codecov](https://codecov.io/gh/juaml/junifer/branch/main/graph/badge.svg?token=5H21JuZXMw)](https://codecov.io/gh/juaml/junifer) + +## About + +junifer is a data handling and feature extraction library targeted towards neuroimaging data specifically functional MRI data. + +It is curently being developed and maintained at the [Applied Machine Learning](https://www.fz-juelich.de/en/inm/inm-7/research-groups/applied-machine-learning-aml) group at [Forschungszentrum Juelich](https://www.fz-juelich.de/en), Germany. Although the library is designed for people working at [Institute of Neuroscience and Medicine - Brain and Behaviour (INM-7)](https://www.fz-juelich.de/en/inm/inm-7), it is designed to be as modular as possible thus enabling others to extend it easily. + +The documentation is available at [https://juaml.github.io/junifer](https://juaml.github.io/junifer/main/index.html). ## Repository Organization * `docs`: Documentation, built using sphinx. * `examples`: Examples, using sphinx-gallery. File names of examples that create visual output must start with `plot_`, otherwise, with `run_`. -* `junifer`: Main library directory - * `api`: User API module - * `data`: Module that handles data required for the library to work (e.g. atlases) +* `junifer`: Main library directory. + * `api`: User API module. + * `configs`: Module for pre-defined configs for most used computing clusters. + * `data`: Module that handles data required for the library to work (e.g. atlases). * `datagrabber`: DataGrabber module. * `datareader`: DataReader module. * `markers`: Markers module. * `pipeline`: Pipeline module. * `preprocess`: Preprocessing module. * `storage`: Storage module. + * `testing`: Testing components module. * `utils`: Utilities module (e.g. logging) - - \ No newline at end of file + +## Installation + +Use `pip` to install from PyPI like so: + +``` +pip install junifer +``` + +## Citation + +If you use junifer in a scientific publication, we would appreciate if you cite our work. Currently, we do not have a publication, so feel free to use the project [URL](https://juaml.github.io/junifer). + +## Contribution + +Contributions are welcome and greatly appreciated. Please read the [guidelines](https://juaml.github.io/junifer/main/contributing.html) to get started. + +## License + +junifer is released under the AGPL v3 license: + +julearn, FZJuelich AML neuroimaging feature extraction library. +Copyright (C) 2022, authors of junifer. + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as published by +the Free Software Foundation, either version 3 of the License, or any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . diff --git a/conda-env.yml b/conda-env.yml new file mode 100644 index 000000000..2801c724a --- /dev/null +++ b/conda-env.yml @@ -0,0 +1,32 @@ +name: junifer-dev +channels: + - conda-forge + - defaults +dependencies: + - python=3.10 + - click>=8.1.3,<8.2 + - numpy>=1.22,<1.23 + - datalad>=0.15.4,<0.18 + - pandas>=1.4.0,<1.5 + - nibabel>=3.2.0,<4.1 + - nilearn>=0.9.0,<1.0 + - sqlalchemy>=1.4.27,<= 1.5.0 + - pyyaml>=5.1.2,<7.0 + - seaborn>=0.11.2,<0.12 + - Sphinx>=5.0.2,<5.1 + - sphinx-gallery>=0.10.1,<0.11 + - numpydoc>=1.4.0,<1.5 + - tox + - ipykernel + - isort + - pytest-cov + - pytest + - black + - flake8 + - flake8-docstrings + - flake8-bugbear + - codespell + - pip + - pip: + - sphinx-rtd-theme>=1.0.0,<1.1 + - sphinx-multiversion>=0.2.4,<0.3 diff --git a/dev-requirements.txt b/dev-requirements.txt deleted file mode 100644 index aad120b7c..000000000 --- a/dev-requirements.txt +++ /dev/null @@ -1,2 +0,0 @@ -flake8 -pytest \ No newline at end of file diff --git a/docs-requirements.txt b/docs-requirements.txt deleted file mode 100644 index 08da25ce4..000000000 --- a/docs-requirements.txt +++ /dev/null @@ -1,6 +0,0 @@ -seaborn -sphinx -sphinx-gallery -sphinx_rtd_theme -git+https://github.com/dls-controls/sphinx-multiversion.git@only-arg -numpydoc \ No newline at end of file diff --git a/docs/api.rst b/docs/api.rst deleted file mode 100644 index 5f50340cb..000000000 --- a/docs/api.rst +++ /dev/null @@ -1,17 +0,0 @@ -# Authors: Federico Raimondo -# License: AGPL -Reference -========= -.. include:: links.inc - -Data Grabbers -^^^^^^^^^^^^^ - -.. autoclass:: junifer.datagrabber.base.BaseDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.BIDSDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.DataladDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.BIDSDataladDataGrabber - :members: \ No newline at end of file diff --git a/docs/api/api.rst b/docs/api/api.rst new file mode 100644 index 000000000..b1a96e243 --- /dev/null +++ b/docs/api/api.rst @@ -0,0 +1,7 @@ + +API Functions +^^^^^^^^^^^^^^ + +.. automodule:: junifer.api + :members: + :imported-members: diff --git a/docs/api/datagrabbers.rst b/docs/api/datagrabbers.rst new file mode 100644 index 000000000..974bd303b --- /dev/null +++ b/docs/api/datagrabbers.rst @@ -0,0 +1,6 @@ +Data Grabbers +^^^^^^^^^^^^^ + +.. automodule:: junifer.datagrabber + :members: + :imported-members: diff --git a/docs/api/datareaders.rst b/docs/api/datareaders.rst new file mode 100644 index 000000000..3dee2aa11 --- /dev/null +++ b/docs/api/datareaders.rst @@ -0,0 +1,7 @@ + +Data Readers +^^^^^^^^^^^^ + +.. automodule:: junifer.datareader + :members: + :imported-members: diff --git a/docs/api/index.rst b/docs/api/index.rst new file mode 100644 index 000000000..4a16b9f78 --- /dev/null +++ b/docs/api/index.rst @@ -0,0 +1,27 @@ +API Reference +============= + +Pipeline Elements +^^^^^^^^^^^^^^^^^ + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + datagrabbers + datareaders + preprocessing + markers + storage + + +Utilities +^^^^^^^^^ + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + api + utils + testing \ No newline at end of file diff --git a/docs/api/markers.rst b/docs/api/markers.rst new file mode 100644 index 000000000..7c6a34b8d --- /dev/null +++ b/docs/api/markers.rst @@ -0,0 +1,7 @@ + +Markers +^^^^^^^ + +.. automodule:: junifer.markers + :members: + :imported-members: diff --git a/docs/api/preprocessing.rst b/docs/api/preprocessing.rst new file mode 100644 index 000000000..571bc7b6e --- /dev/null +++ b/docs/api/preprocessing.rst @@ -0,0 +1,7 @@ + +Pre-processing +^^^^^^^^^^^^^^ + +.. automodule:: junifer.preprocess + :members: + :imported-members: diff --git a/docs/api/storage.rst b/docs/api/storage.rst new file mode 100644 index 000000000..4e785ce60 --- /dev/null +++ b/docs/api/storage.rst @@ -0,0 +1,7 @@ + +Storage +^^^^^^^ + +.. automodule:: junifer.storage + :members: + :imported-members: diff --git a/docs/api/testing.rst b/docs/api/testing.rst new file mode 100644 index 000000000..0220bd917 --- /dev/null +++ b/docs/api/testing.rst @@ -0,0 +1,6 @@ + +Testing +^^^^^^^ + +.. automodule:: junifer.testing.datagrabbers + :members: diff --git a/docs/api/utils.rst b/docs/api/utils.rst new file mode 100644 index 000000000..999a180d3 --- /dev/null +++ b/docs/api/utils.rst @@ -0,0 +1,7 @@ + +Utils +^^^^^ + +.. automodule:: junifer.utils + :members: + :imported-members: diff --git a/docs/builtin.rst b/docs/builtin.rst new file mode 100644 index 000000000..e5ece4e49 --- /dev/null +++ b/docs/builtin.rst @@ -0,0 +1,112 @@ + +Available pipeline steps +======================== + + +Data Grabbers +^^^^^^^^^^^^^ + +.. + Provide a list of the DataGrabbers that are implemented or planned. + Access: Valid options are + - Open + - Open with registration + - Restricted + + Type/config: this should mention weather the class is built-in in the + core of junifer or needs to be imported from a specific configuration in + the `junifer.configs` module. + + State: this should indicate the state of the dataset. Valid options are + - Planned + - In Progress + - Done + + Version added: If the status is "Done", the Junifer version in which the + dataset was added. Else, a link to the Github issue or pull request + implementing the dataset. Links to github can be added by using the + following syntax: :gh:`` + +.. list-table:: Available data grabbers + :widths: auto + :header-rows: 1 + + * - Class + - Description + - Access + - Type/Config + - State + - Version Added + * - `DataladHCP1200` + - `HCP OpenAccess dataset `_ + - Open with registration + - Built-in + - In Progress + - :gh:`4` + * - `JuselessDataladUKBVBM` + - UKB VBM dataset preprocessed with CAT. Available for Juseless only + - Restricted + - `junifer.configs.juseless` + - Done + - 0.0.1 + + + +Markers +^^^^^^^ + +.. + Provide a list of the Markers that are implemented or planned. + + State: this should indicate the state of the dataset. Valid options are + - Planned + - In Progress + - Done + + Version added: If the status is "Done", the Junifer version in which the + dataset was added. Else, a link to the Github issue or pull request + implementing the dataset. Links to github can be added by using the + following syntax: :gh:`` + +.. list-table:: Available data grabbers + :widths: auto + :header-rows: 1 + + * - Class + - Description + - State + - Version Added + * - :class:`junifer.markers.ParcelAggregation` + - Apply parcellation and perform aggregation function + - Done + - 0.0.1 + + + +Available Atlases and Coordinates +================================= + ++------------------+-----------------------+-----------------------------+---------------+ +| Name | Options | Keys | Version Added | ++==================+=======================+=============================+===============+ +| Schaefer | `n_rois` | `Schaefer100x7` | 0.0.1 | +| | `yeo_networks` | `Schaefer200x7` | | +| | | `Schaefer300x7` | | +| | | `Schaefer400x7` | | +| | | `Schaefer500x7` | | +| | | `Schaefer600x7` | | +| | | `Schaefer700x7` | | +| | | `Schaefer800x7` | | +| | | `Schaefer900x7` | | +| | | `Schaefer1000x7` | | +| | | `Schaefer100x17` | | +| | | `Schaefer200x17` | | +| | | `Schaefer300x17` | | +| | | `Schaefer400x17` | | +| | | `Schaefer500x17` | | +| | | `Schaefer600x17` | | +| | | `Schaefer700x17` | | +| | | `Schaefer800x17` | | +| | | `Schaefer900x17` | | +| | | `Schaefer1000x17` | | ++------------------+-----------------------+-----------------------------+---------------+ diff --git a/docs/changes/contributors.inc b/docs/changes/contributors.inc index 50ed21edf..ff1c379f9 100644 --- a/docs/changes/contributors.inc +++ b/docs/changes/contributors.inc @@ -1,4 +1,5 @@ .. _Fede Raimondo: https://fraimondo.github.io .. _Kaustubh Patil: https://github.com/kaurao .. _Leonard Sasse: https://github.com/LeSasse -.. _Amir Omidvarnia: https://github.com/omidvarnia \ No newline at end of file +.. _Amir Omidvarnia: https://github.com/omidvarnia +.. _Synchon Mandal: https://github.com/synchon \ No newline at end of file diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index 7c826721b..2a6627ca6 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -8,6 +8,10 @@ - "Bugs" for bug fixes - "API changes" for backward-incompatible changes +.. NOTE: add the contributors and reference to the github issue/PR at the end + Example: + - Implemented feature X (:gh:`151` by `Sami Hamdan`_). + .. _current: Current (0.0.0.dev) diff --git a/docs/contribution.rst b/docs/contribution.rst new file mode 100644 index 000000000..606341056 --- /dev/null +++ b/docs/contribution.rst @@ -0,0 +1,217 @@ +.. include:: links.inc + +Contributing to junifer +======================= + + +Setting up the local development environment +-------------------------------------------- + +1. Fork the https://github.com/juaml/junifer repository on GitHub. If you + have never done this before, `follow the official guide + `_. +2. Clone your fork locally as described in the same guide. +3. Install your local copy into a Python virtual environment. You can `read + this guide to learn more + `_ about them + and how to create one. + + .. code-block:: console + + pip install -e ".[dev]" + +4. Create a branch for local development using the ``main`` branch as a + starting point. Use ``fix``, ``refactor``, or ``feat`` as a prefix. + + .. code-block:: console + + git checkout dev + git checkout -b / + + Now you can make your changes locally. + +5. When making changes locally, it is helpful to ``git commit`` your work + regularly. On one hand to save your work and on the other hand, the smaller + the steps, the easier it is to review your work later. Please use `semantic + commit messages + `_. + + .. code-block:: console + + git add . + git commit -m ": " + +6. When you're done making changes, check that your changes pass our test suite. + This is all included with ``tox``. + + .. code-block:: console + + tox + + You can also run all ``tox`` tests in parallel. As of ``tox 3.7``, you can run + + .. code-block:: console + + tox --parallel + + +7. Push your branch to GitHub. + + .. code-block:: console + + git push origin / + +8. Open the link displayed in the message when pushing your new branch in order + to submit a pull request. Please follow the template presented to you in the + web interface to complete your pull request. + + +GitHub Pull Request guidelines +------------------------------ + +Before you submit a pull request, check that it meets these guidelines: + +1. The pull request should include tests in the respective ``tests`` directory. + Except in rare circumstances, code coverage must not decrease (as reported + by codecov which runs automatically when you submit your pull request). +2. If the pull request adds functionality, the docs should be + updated. Consider creating a Python file that demonstrates the usage in + ``examples/`` directory. +3. The pull request should also include a short one-liner of your contribution + in ``docs/changes/latest.inc``. If it's your first contribution, also add + yourself to ``docs/changes/contributors.inc``. +4. The pull request will be tested against several Python versions. +5. Someone from the core team will review your work and guide you to a successful + contribution. + + +Running unit tests +------------------ + +junifer uses `pytest `_ for its +unit-tests and new features should in general always come with new +tests that make sure that the code runs as intended. + +To run all tests + +.. code-block:: console + + tox -e test + + +Adding and building documentation +--------------------------------- + +Building the documentation requires some extra packages and can be installed by + +.. code-block:: console + + pip install -e ".[docs]" + +To build the docs + +.. code-block:: bash + + cd docs + make local + +To view the documentation, open ``docs/_build/html/index.html``. + +In case you remove some files or change their filenames, you can run into +errors when using ``make local``. In this situation you can use ``make clean`` +to clean up the already build files and then re-run ``make local``. + + +Writing Examples +---------------- + +The format used for text is reST. Check the `sphinx reST reference`_ for more +details. The examples are run and displayed in HTML format using `sphinx gallery`_. To add an +example, just create a ``.py`` file that starts either with ``plot_`` or ``run_``, +dependending on whether the example generates a figure or not. + +The first lines of the example should be a Python block comment with a title, +a description of the example, authors and license name. + +The following is an example of how to start an example + +.. code-block:: python + + """ + Generic BIDS datagrabber for datalad. + ===================================== + + This example uses a generic BIDS datagraber to get the data from a BIDS dataset + store in a datalad remote sibling. + + Authors: Federico Raimondo + + License: BSD 3 clause + """ + +The rest of the script will be executed as normal Python code. In order to +render the output and embed formatted text within the code, you need to add +a 79 ``#`` (a full line) at the point in which you want to render and add text. +Each line of text shall be preceded with ``#``. The code that is not +commented will be executed. + +The following example will create texts and render the output between the +texts. + +.. code-block:: python + + from junifer.datagrabber import PatternDataladDataGrabber + from junifer.utils import configure_logging + + + ############################################################################### + # Set the logging level to info to see extra information + configure_logging(level="INFO") + + + ############################################################################### + # The BIDS datagrabber requires three parameters: the types of data we want, + # the specific pattern that matches each type, and the variables that will be + # replaced int he patterns. + types = ["T1w", "bold"] + patterns = { + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", + } + replacements = ["subject"] + ############################################################################### + # Additionally, a datalad datagrabber requires the URI of the remote sibling + # and the location of the dataset within the remote sibling. + repo_uri = "https://gin.g-node.org/juaml/datalad-example-bids" + rootdir = "example_bids" + + ############################################################################### + # Now we can use the datagrabber within a `with` context + # One thing we can do with any datagrabber is iterate over the elements. + # In this case, each element of the datagrabber is one session. + with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, + ) as dg: + for elem in dg: + print(elem) + + ############################################################################### + # Another feature of the datagrabber is the ability to get a specific + # element by its name. In this case, we index `sub-01` and we get the file + # paths for the two types of data we want (T1w and bold). + with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, + ) as dg: + sub01 = dg["sub-01"] + print(sub01) + +Finally, when the example is done, you can run it as a normal Python script. +To generate the HTML, just build the docs. diff --git a/docs/faq.rst b/docs/faq.rst new file mode 100644 index 000000000..7e49b63a5 --- /dev/null +++ b/docs/faq.rst @@ -0,0 +1,4 @@ +.. include:: links.inc + +FAQs +==== diff --git a/docs/index.rst b/docs/index.rst index 1f7f0d2aa..cef0f0606 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -1,17 +1,31 @@ .. include:: links.inc -Welcome to the documentation! -============================= +Welcome to junifer's documentation! +=================================== +junifer (JUelich NeuroImaging FEature extractoR) is a data handling and feature +extraction library targeted towards neuroimaging data specifically functional +MRI data. + +It is curently being developed and maintained at the Applied Machine Learning +(`AML`_) group at Forschungszentrum Juelich, Germany. Although the library is +designed for people working at Institute of Neuroscience and Medicine - Brain +and Behaviour (`INM-7`_), it is designed to be as modular as possible thus +enabling others to extend it easily. .. toctree:: + :numbered: :maxdepth: 2 :caption: Contents: installation - api + understanding/index.rst + builtin auto_examples/index.rst + api/index.rst + contribution maintaining + faq whats_new @@ -21,4 +35,3 @@ Indices and tables * :ref:`genindex` * :ref:`modindex` * :ref:`search` - diff --git a/docs/installation.rst b/docs/installation.rst index 99ddee684..90209b8e5 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -1,91 +1,53 @@ .. include:: links.inc -Installing -========== +Installing junifer +================== Requirements ^^^^^^^^^^^^ -junifer requires the following packages: +junifer is compatible with `Python`_ >= 3.8 and requires the following packages: -Running the examples requires: +* click>=8.1.3,<8.2 +* numpy>=1.22,<1.23 +* datalad>=0.15.4,<0.18 +* pandas>=1.4.0,<1.5 +* nibabel>=3.2.0,<4.1 +* nilearn>=0.9.0,<1.0 +* sqlalchemy>=1.4.27,<= 1.5.0 +* pyyaml>=5.1.2,<7.0 -Depending on the installation method, this packages might be installed -automatically. +Depending on the installation method, these packages might be installed automatically. -Installing -^^^^^^^^^^ -There are different ways to install junifer: +Installation +^^^^^^^^^^^^ +Depending on your use-case, junifer can be installed differently: * Install the :ref:`install_latest_release`. This is the most suitable approach - for most end users. -* Install the :ref:`install_latest_development`. This version will have the - latest features. However, it is still under development and not yet - officially released. Some features might still change before the next stable - release. -* Install from :ref:`install_development_git`. This is mostly suitable for - developers that want to have the latest version and yet edit the code. + for end users. +* Install from :ref:`install_development_git`. This is the most suitable approach + for developers. -Either way, we strongly recommend using virtual environments: - -* `venv`_ -* `conda env`_ +Either way, we strongly recommend using `virtual environments `_. .. _install_latest_release: -Latest release +Stable release -------------- -We have packaged junifer and published it in PyPi, so you can just install it -with `pip`. +Use ``pip`` to install julearn from `PyPI `_, like so: .. code-block:: bash - pip install -U junifer - - -.. _install_latest_development: - -Latest Development Version --------------------------- -First, make sure that you have all the dependencies installed: - -Then, install junifer from TestPypi - -.. code-block:: bash - - pip install -U junifer --pre + pip install junifer .. _install_development_git: -Local git repository (for developers) -------------------------------------- -First, make sure that you have all the dependencies installed: +Local Git repository +-------------------- -Then, clone `junifer Github`_ repository in a folder of your choice: - -.. code-block:: bash - - git clone https://github.com/juaml/junifer.git - -Install development mode requirements: - -.. code-block:: bash - - cd junifer - pip install -r dev-requirements.txt - -Finally, install in development mode: - -.. code-block:: bash - - python setup.py develop - -.. note:: Every time that you run ``setup.py develop``, the version is going to - be automatically set based on the git history. Nevertheless, this change - should not be committed (changes to ``_version.py``). Running ``git stash`` - at this point will forget the local changes to ``_version.py``. +Follow the `detailed contribution guidelines `_. diff --git a/docs/links.inc b/docs/links.inc index 27414754b..47d3ad6a4 100644 --- a/docs/links.inc +++ b/docs/links.inc @@ -11,6 +11,7 @@ .. _`AML`: https://www.fz-juelich.de/inm/inm-7/EN/Forschung/Applied%20Machine%20Learning/_node.html .. _`INM-7`: https://www.fz-juelich.de/inm/inm-7/EN/Home/home_node.html +.. _`julearn`: https://juaml.github.io/julearn .. _`pandas`: https://pandas.pydata.org .. _`pandas.DataFrame` : https://pandas.pydata.org/pandas-docs/stable/reference/api/pandas.DataFrame.html @@ -30,4 +31,4 @@ .. _`setuptools_scm`: https://github.com/pypa/setuptools_scm/ .. _`sphinx gallery`: https://sphinx-gallery.github.io/stable/index.html -.. _`sphinx RST reference`: https://www.sphinx-doc.org/en/master/usage/restructuredtext/basics.html#inline-markup \ No newline at end of file +.. _`sphinx reST reference`: https://www.sphinx-doc.org/en/master/usage/restructuredtext/basics.html#inline-markup diff --git a/docs/sphinxext/gh_substitutions.py b/docs/sphinxext/gh_substitutions.py index 47b5e86f6..e9e651fd9 100644 --- a/docs/sphinxext/gh_substitutions.py +++ b/docs/sphinxext/gh_substitutions.py @@ -19,7 +19,7 @@ def gh_role(name, rawtext, text, lineno, inliner, options={}, content=[]): else: slug = 'issues/' + text text = '#' + text - ref = 'https://github.com/juaml/julearn/' + slug + ref = 'https://github.com/juaml/junifer/' + slug set_classes(options) node = reference(rawtext, text, refuri=ref, **options) return [node], [] diff --git a/docs/understanding/data.rst b/docs/understanding/data.rst new file mode 100644 index 000000000..1010720a8 --- /dev/null +++ b/docs/understanding/data.rst @@ -0,0 +1,46 @@ +.. include:: ../links.inc + +The Data Object +=============== + +Description +^^^^^^^^^^^ + +This is the *object* that traverses the steps of the pipeline. It is indeed a +dictionary of dictionaries. The first level of keys are the :ref:`data_types` +and a special key named ``meta`` that contains all the information on the data +object including source and previous transformation steps. + +The second level of keys are the actual data. So far, there are two keys used: + +- ``path``: path to the file containing the data. +- ``data``: the data loaded in memory. + +The :ref:`datagrabber` step will only fill the ``path`` value. +The ``data`` value will be filled by the :ref:`datareader` step, if it is one of the possible file types +that the datareader can read. + +.. _data_types: + +Data types +^^^^^^^^^^ + +.. list-table:: Built-in data types + :widths: 30 80 40 + :header-rows: 1 + + * - Name + - Description + - Example + * - ``T1w`` + - T1w image (3D) + - Preprocessed or Raw T1w image + * - ``BOLD`` + - BOLD image (4D) + - Preprocessed/Denoised BOLD image (fmriprep output) + * - ``VBM_GM`` + - VBM Gray Matter segmentation (3D) + - CAT output (`m0wp1` images) + * - ``VBM_WM`` + - VBM White Matter segmentation (3D) + - CAT output (`m0wp2` images) diff --git a/docs/understanding/datagrabber.rst b/docs/understanding/datagrabber.rst new file mode 100644 index 000000000..1a23ed3fa --- /dev/null +++ b/docs/understanding/datagrabber.rst @@ -0,0 +1,6 @@ +.. include:: ../links.inc + +.. _datagrabber: + +Data Grabber +============ diff --git a/docs/understanding/datareader.rst b/docs/understanding/datareader.rst new file mode 100644 index 000000000..8a3521efd --- /dev/null +++ b/docs/understanding/datareader.rst @@ -0,0 +1,6 @@ +.. include:: ../links.inc + +.. _datareader: + +Data Reader +=========== diff --git a/docs/understanding/index.rst b/docs/understanding/index.rst new file mode 100644 index 000000000..e24c92883 --- /dev/null +++ b/docs/understanding/index.rst @@ -0,0 +1,28 @@ +.. include:: ../links.inc + +Understanding junifer +===================== + +Before you start, you should understand how junifer works. Junifer is a +tool conceived to extract features from neuroimaging data in a easy-to-use +manner, with minimal coding and minimal user expertise in the internal aspects. + +Unlike other tools like FSL, SPM, AFNI, etc., junifer is not a toolbox to +pre-process data, but a toolbox to extract features from previously pre-processed +data. + +The main idea is that you have a set of images (e.g. a set of functional MRI, +structural MRI, diffusion MRI, etc.) and you want to extract features to +later use in stastical analyses or machine learning (for example, using +julearn_). + + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + data + datagrabber + datareader + marker + storage diff --git a/docs/understanding/marker.rst b/docs/understanding/marker.rst new file mode 100644 index 000000000..3bf018ba7 --- /dev/null +++ b/docs/understanding/marker.rst @@ -0,0 +1,4 @@ +.. include:: ../links.inc + +Marker +====== diff --git a/docs/understanding/storage.rst b/docs/understanding/storage.rst new file mode 100644 index 000000000..ef296e1d4 --- /dev/null +++ b/docs/understanding/storage.rst @@ -0,0 +1,4 @@ +.. include:: ../links.inc + +Storage +======= diff --git a/examples/norun_hcpfc_pearson.py b/examples/norun_hcpfc_pearson.py index d2d17640d..d6a460440 100644 --- a/examples/norun_hcpfc_pearson.py +++ b/examples/norun_hcpfc_pearson.py @@ -1,65 +1,81 @@ """ HCP FC Extraction ====================== + Authors: Leonard Sasse License: BSD 3 clause """ -from junifer.api import run_pipeline +from junifer.api import run + + +datagrabber = { + "kind": "HCPOpenAccess", + "modality": "fMRI", + "preprocessed": "ICA+FIX", + "space": "volumetric", +} custom_confound_strategy = { - 'filter': 'butterworth', - 'detrend': True, - 'high_pass': 0.01, - 'low_pass': 0.08, - 'standardize': True, - 'confounds': ['csf', 'wm', 'gsr'], - 'derivatives': True, - 'squares': True, - 'other': [] + "filter": "butterworth", + "detrend": True, + "high_pass": 0.01, + "low_pass": 0.08, + "standardize": True, + "confounds": ["csf", "wm", "gsr"], + "derivatives": True, + "squares": True, + "other": [], } markers = [ - {'name': 'Power264_FCPearson', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Power264', - 'method': 'Pearson', - 'confound_strategy': 'Params36'}, - {'name': 'Schaefer400x17_FCPearson', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Pearson', - 'confound_strategy': 'Params24'}, - {'name': 'Power264_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Power264', - 'method': 'Spearman', - 'confound_strategy': 'ICAAROMA'}, - {'name': 'Schaefer400x17_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Spearman', - 'confound_strategy': 'path/to/predefined/confound_file.tsv'}, - {'name': 'Schaefer400x17_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Spearman', - 'confound_strategy': custom_confound_strategy} + { + "name": "Power264_FCPearson", + "kind": "FunctionalConnectivity", + "atlas": "Power264", + "method": "Pearson", + "confound_strategy": "Params36", + }, + { + "name": "Schaefer400x17_FCPearson", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Pearson", + "confound_strategy": "Params24", + }, + { + "name": "Power264_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Power264", + "method": "Spearman", + "confound_strategy": "ICAAROMA", + }, + { + "name": "Schaefer400x17_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Spearman", + "confound_strategy": "path/to/predefined/confound_file.tsv", + }, + { + "name": "Schaefer400x17_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Spearman", + "confound_strategy": custom_confound_strategy, + }, ] -dg_params = { - 'modality': 'fMRI', - 'preprocessed': 'ICA+FIX', - 'space': 'volumetric', +storage = { + "kind": "SQLiteFeatureStorage", + "uri": "/data/project/juniferexample", } -run_pipeline( - workdir='/tmp', - datagrabber='HCPOpenAccess', - datagrabber_params=dg_params, - element=('100408', 'REST1', "LR"), +run( + workdir="/tmp", + datagrabber=datagrabber, + elements=[("100408", "REST1", "LR")], markers=markers, - storage='SQLDataFrameStorage', - storage_params={'outpath': '/data/project/juniferexample'}, + storage=storage, ) diff --git a/examples/norun_ukbvm_gmd.py b/examples/norun_ukbvm_gmd.py index 168478d4c..dc4ca5c4a 100644 --- a/examples/norun_ukbvm_gmd.py +++ b/examples/norun_ukbvm_gmd.py @@ -8,28 +8,36 @@ License: BSD 3 clause """ -from junifer.api import run_pipeline +from junifer.api import run + markers = [ - {'name': 'Schaefer1000x7_TrimMean80', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'trimmean80'}, - {'name': 'Schaefer1000x7_Mean', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'mean'}, - {'name': 'Schaefer1000x7_Std', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'std'} + { + "name": "Schaefer1000x7_TrimMean80", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "trim_mean", + "method_params": {"proportiontocut": 0.2}, + }, + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, ] -run_pipeline( - workdir='/tmp', - datagrabber='JuselessUKBVBM', - element=('sub-1627474', 'ses-2'), +run( + workdir="/tmp", + datagrabber="JuselessUKBVBM", + elements=("sub-1627474", "ses-2"), markers=markers, - storage='SQLDataFrameStorage', - storage_params={'outpath': '/data/project/juniferexample'}, + storage="SQLDataFrameStorage", + storage_params={"outpath": "/data/project/juniferexample"}, ) diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py new file mode 100644 index 000000000..1f76a88f2 --- /dev/null +++ b/examples/run_compute_parcel_mean.py @@ -0,0 +1,57 @@ +""" +Computer Parcel Aggregation. +============================ + +This example uses a ParcelAggregation marker to compute the mean of each parcel +using the Schaefer atlas (100 rois, 7 Yeo networks) for both a 3D and 4D nifti + +Authors: Federico Raimondo + +License: BSD 3 clause +""" + +import nilearn + +from junifer.markers.parcel import ParcelAggregation +from junifer.utils import configure_logging + + +############################################################################### +# Set the logging level to info to see extra information +configure_logging(level="INFO") + + +############################################################################### +# Load the VBM GM data (3d): +# - Fetch the Oasis dataset +oasis_dataset = nilearn.datasets.fetch_oasis_vbm(n_subjects=1) +vbm_fname = oasis_dataset.gray_matter_maps[0] +vbm_img = nilearn.image.load_img(vbm_fname) + +############################################################################### +# Load the functional data (4d): +# - Fetch the SPM auditory dataset +# - Concatenate the functional data into one 4D image +s_func_data = nilearn.datasets.fetch_spm_auditory() +fmri_img = nilearn.image.concat_imgs(s_func_data.func) + +############################################################################### +# Define the marker +marker = ParcelAggregation(atlas="Schaefer100x7", method="mean") + +############################################################################### +# Prepare the input +input = {"BOLD": {"data": fmri_img}, "VBM_GM": {"data": vbm_img}} + +############################################################################### +# Fit transform the data +out = marker.fit_transform(input) + +############################################################################### +# Check the results + +print(out.keys()) +print(out["VBM_GM"]["data"].shape) # Shape is (1 x parcels) + +print(out.keys()) +print(out["BOLD"]["data"].shape) # Shape is (timepoints x parcels) diff --git a/examples/run_datagrabber_bids_datalad.py b/examples/run_datagrabber_bids_datalad.py index 221b4f74c..eb14892a6 100644 --- a/examples/run_datagrabber_bids_datalad.py +++ b/examples/run_datagrabber_bids_datalad.py @@ -10,35 +10,42 @@ Authors: Federico Raimondo License: BSD 3 clause """ -from junifer.datagrabber.base import BIDSDataladDataGrabber +from junifer.datagrabber import PatternDataladDataGrabber from junifer.utils import configure_logging + ############################################################################### # Set the logging level to info to see extra information -configure_logging(level='INFO') +configure_logging(level="INFO") ############################################################################### -# The BIDS datagrabber requires two parameters: the types of data we want, -# and the specific pattern that matches each type. -types = ['T1w', 'bold'] +# The BIDS datagrabber requires three parameters: the types of data we want, +# the specific pattern that matches each type, and the variables that will be +# replaced int he patterns. +types = ["T1w", "bold"] patterns = { - 'T1w': 'anat/{subject}_T1w.nii.gz', - 'bold': 'func/{subject}_task-rest_bold.nii.gz' + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", } - +replacements = ["subject"] ############################################################################### # Additionally, a datalad datagrabber requires the URI of the remote sibling # and the location of the dataset within the remote sibling. -repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids' -rootdir = 'example_bids' +repo_uri = "https://gin.g-node.org/juaml/datalad-example-bids" +rootdir = "example_bids" ############################################################################### # Now we can use the datagrabber within a `with` context # One thing we can do with any datagrabber is iterate over the elements. # In this case, each element of the datagrabber is one session. -with BIDSDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri) as dg: +with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, +) as dg: for elem in dg: print(elem) @@ -46,7 +53,12 @@ with BIDSDataladDataGrabber(rootdir=rootdir, types=types, # Another feature of the datagrabber is the ability to get a specific # element by its name. In this case, we index `sub-01` and we get the file # paths for the two types of data we want (T1w and bold). -with BIDSDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri) as dg: - sub01 = dg['sub-01'] +with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, +) as dg: + sub01 = dg["sub-01"] print(sub01) diff --git a/examples/run_run_gmd_mean.py b/examples/run_run_gmd_mean.py new file mode 100644 index 000000000..2b2ee79f8 --- /dev/null +++ b/examples/run_run_gmd_mean.py @@ -0,0 +1,53 @@ +""" +UKB VBM GMD Extraction +====================== + +Authors: Federico Raimondo + +License: BSD 3 clause +""" +import tempfile + +import junifer.testing.registry # noqa: F401 +from junifer.api import run + + +datagrabber = { + "kind": "OasisVBMTestingDatagrabber", +} + +markers = [ + { + "name": "Schaefer1000x7_TrimMean80", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "trim_mean", + "method_params": {"proportiontocut": 0.2}, + }, + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, +] + +storage = { + "kind": "SQLiteFeatureStorage", +} + +with tempfile.TemporaryDirectory() as tmpdir: + uri = f"{tmpdir}/test.db" + storage["uri"] = uri + run( + workdir="/tmp", + datagrabber=datagrabber, + markers=markers, + storage=storage, + ) diff --git a/examples/yamls/gmd_mean.yaml b/examples/yamls/gmd_mean.yaml new file mode 100644 index 000000000..ce50d5d5c --- /dev/null +++ b/examples/yamls/gmd_mean.yaml @@ -0,0 +1,25 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +elements: +markers: + - name: Schaefer1000x7_TrimMean80 + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: trim_mean + method_params: + proportiontocut: 0.2 + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean + - name: Schaefer1000x7_Std + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: std +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db + diff --git a/examples/yamls/gmd_mean_htcondor.yaml b/examples/yamls/gmd_mean_htcondor.yaml new file mode 100644 index 000000000..fde41b733 --- /dev/null +++ b/examples/yamls/gmd_mean_htcondor.yaml @@ -0,0 +1,20 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /data/group/appliedml/fraimondo/junifer_test/test.db +queue: + jobname: TestHTCondorQueue + kind: HTCondor + env: + kind: conda + name: junifer + mem: 8G \ No newline at end of file diff --git a/examples/yamls/ukb_gmd_mean.yaml b/examples/yamls/ukb_gmd_mean.yaml new file mode 100644 index 000000000..7d2303d7e --- /dev/null +++ b/examples/yamls/ukb_gmd_mean.yaml @@ -0,0 +1,25 @@ +with: junifer.configs.juseless +workdir: /tmp + +datagrabber: + kind: JuselessUKBVBM +elements: +markers: + - name: Schaefer1000x7_TrimMean80 + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: trim_mean + method_params: + proportiontocut: 0.2 + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean + - name: Schaefer1000x7_Std + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: std +storage: + kind: SQLiteFeatureStorage + uri: /data/project/ukb_motor/junifer_test/test.db + diff --git a/junifer/__init__.py b/junifer/__init__.py index ac53f2bcb..5d46af23a 100644 --- a/junifer/__init__.py +++ b/junifer/__init__.py @@ -1,6 +1,17 @@ -from . _version import __version__ +"""Provide imports for junifer package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from ._version import __version__ from . import api -from . import utils +from . import configs +from . import data from . import datagrabber +from . import datareader from . import markers -from . import configs \ No newline at end of file +from . import pipeline +from . import preprocess +from . import storage +from . import utils diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index 51642c08a..9961e2a4a 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -1 +1,8 @@ -from . pipeline import run_pipeline \ No newline at end of file +"""Provide imports for api sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .cli import cli +from .functions import run, collect diff --git a/junifer/api/cli.py b/junifer/api/cli.py new file mode 100644 index 000000000..2402de4b9 --- /dev/null +++ b/junifer/api/cli.py @@ -0,0 +1,195 @@ +"""Provide functions for cli.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import pathlib +from typing import Dict, List, Union + +import click + +from ..utils.logging import configure_logging, logger, warn_with_log +from .functions import collect as api_collect +from .functions import queue as api_queue +from .functions import run as api_run +from .parser import parse_yaml + + +def _parse_elements(element: str, config: Dict) -> Union[List, None]: + """Parse elements from cli. + + Parameters + ---------- + element : str + The element to operate on. + config : dict + The configuration to operate using. + + Returns + ------- + list + The element(s) as list. + + """ + logger.debug(f"Parsing elements: {element}") + if len(element) == 0: + return None + # TODO: If len == 1, check if its a file, then parse elements from file + elements = [x.split(",") if "," in x else x for x in element] + logger.debug(f"Parsed elements: {elements}") + if elements is not None and "elements" in config: + warn_with_log( + "One or more elements have been specified in both the command " + "line and in the config file. The command line has precedence " + "over the configuration file. That is, the elements specified " + "in the command line will be used. The elements specified in " + "the configuration file will be ignored. To remove this warning, " + 'please remove the "elements" item from the configuration file.' + ) + elif elements is None: + elements = config.get("elements", None) + return elements + + +@click.group() +def cli() -> None: # pragma: no cover + """CLI for JUelich NeuroImaging FEature extractoR.""" + + +@cli.command() +@click.argument( + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), +) +@click.option("--element", type=str, multiple=True) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def run(filepath: click.Path, element: str, verbose: click.Choice) -> None: + """Run command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + element : str + The element to operate using. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ + configure_logging(level=str(verbose).upper()) + # TODO: add validation + config = parse_yaml(filepath) # type: ignore + workdir = config["workdir"] + datagrabber = config["datagrabber"] + markers = config["markers"] + storage = config["storage"] + elements = _parse_elements(element, config) + # Perform operation + api_run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=elements, + ) + + +@cli.command() +@click.argument( + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), +) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def collect(filepath: click.Path, verbose: click.Choice) -> None: + """Collect command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ + configure_logging(level=str(verbose).upper()) + # TODO: add validation + config = parse_yaml(filepath) # type: ignore + storage = config["storage"] + # Perform operation + api_collect(storage=storage) + + +@cli.command() +@click.argument( + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), +) +@click.option("--element", type=str, multiple=True) +@click.option("--overwrite", is_flag=True) +@click.option("--submit", is_flag=True) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def queue( + filepath: click.Path, + element: str, + overwrite: bool, + submit: bool, + verbose: click.Choice, +) -> None: + """Queue command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + element : str + The element to operate using. + overwrite : bool + Whether to overwrite existing directory. + submit : bool + Whether to submit the job. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ + configure_logging(level=str(verbose).upper()) + # TODO: add validation + config = parse_yaml(filepath) # type: ignore + elements = _parse_elements(element, config) + queue_config = config.pop("queue") + kind = queue_config.pop("kind") + api_queue( + config=config, + kind=kind, + overwrite=overwrite, + elements=elements, + submit=submit, + **queue_config, + ) + + +@cli.command() +def selftest() -> None: + """Selftest command for CLI.""" + pass diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index fa8ac0257..49011871d 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -1,13 +1,17 @@ +"""Provide decorators for api.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from . pipeline import register + +from .registry import register -def register_datagrabber(klass): - """Datagrabber decorator. +def register_datagrabber(klass: type) -> type: + """Datagrabber registration decorator. - Registers the datagrabber so it can be used by name in the pipeline. + Registers the datagrabber so it can be used by name. Parameters ---------- @@ -17,7 +21,60 @@ def register_datagrabber(klass): Returns ------- klass: class - The unmodified input class + The unmodified input class. + """ - register('datagrabber', klass.__name__, klass) + register( + step="datagrabber", + name=klass.__name__, + klass=klass, + ) + return klass + + +def register_marker(klass: type) -> type: + """Marker registration decorator. + + Registers the marker so it can be used by name. + + Parameters + ---------- + klass: class + The class of the marker to register. + + Returns + ------- + klass: class + The unmodified input class. + + """ + register( + step="marker", + name=klass.__name__, + klass=klass, + ) + return klass + + +def register_storage(klass): + """Storage registration decorator. + + Registers the storage so it can be used by name. + + Parameters + ---------- + klass: class + The class of the storage to register. + + Returns + ------- + klass: class + The unmodified input class. + + """ + register( + step="storage", + name=klass.__name__, + klass=klass, + ) return klass diff --git a/junifer/api/functions.py b/junifer/api/functions.py new file mode 100644 index 000000000..f8ba24aa8 --- /dev/null +++ b/junifer/api/functions.py @@ -0,0 +1,571 @@ +"""Provide functions for cli.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import shutil +import subprocess +import typing +from pathlib import Path +from typing import Dict, List, Optional, Tuple, Union + +import yaml + +from ..datagrabber.base import BaseDataGrabber +from ..markers.base import BaseMarker +from ..markers.collection import MarkerCollection +from ..storage.base import BaseFeatureStorage +from ..utils import logger, raise_error +from ..utils.fs import make_executable +from .registry import build + + +def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber: + """Get datagrabber. + + Parameters + ---------- + datagrabber_config : dict + The config to get the datagrabber using. + + Returns + ------- + dict + The datagrabber. + + """ + datagrabber_params = datagrabber_config.copy() + datagrabber_kind = datagrabber_params.pop("kind") + datagrabber = build( + step="datagrabber", + name=datagrabber_kind, + baseclass=BaseDataGrabber, + init_params=datagrabber_params, + ) + datagrabber = typing.cast(BaseDataGrabber, datagrabber) + return datagrabber + + +def run( + workdir: Union[str, Path], + datagrabber: Dict, + markers: List[Dict], + storage: Dict, + elements: Union[str, List[Union[str, Tuple]], Tuple, None] = None, +) -> None: + """Run the pipeline on the selected element. + + Parameters + ---------- + workdir : str or pathlib.Path + Directory where the pipeline will be executed. + datagrabber : dict + Datagrabber to use. Must have a key 'kind' with the kind of + datagrabber to use. All other keys are passed to the datagrabber + init function. + markers : list of dict + List of markers to extract. Each marker is a dict with at least two + keys: "name" and "kind". The "name" key is used to name the output + marker. The "kind" key is used to specify the kind of marker to + extract. The rest of the keys are used to pass parameters to the + marker calculation. + storage : dict + Storage to use. Must have a key "kind" with the kind of + storage to use. All other keys are passed to the storage + init function. + elements : str or tuple or list of str or tuple, optional + Element(s) to process. Will be used to index the datagrabber + (default None). + + """ + # Convert str to Path + if isinstance(workdir, str): + workdir = Path(workdir) + if not isinstance(elements, List) and elements is not None: + elements = [elements] + # Get datagrabber to use + datagrabber_object = _get_datagrabber(datagrabber) + # Copy to avoid changing the original dict + _markers = [x.copy() for x in markers] + built_markers = [] + for t_marker in _markers: + kind = t_marker.pop("kind") + t_m = build( + step="marker", + name=kind, + baseclass=BaseMarker, + init_params=t_marker, + ) + built_markers.append(t_m) + # Get storage engine to use + storage_params = storage.copy() + storage_kind = storage_params.pop("kind") + storage_object = build( + step="storage", + name=storage_kind, + baseclass=BaseFeatureStorage, + init_params=storage_params, + ) + storage_object = typing.cast(BaseFeatureStorage, storage_object) + # Create new marker collection + mc = MarkerCollection(markers=built_markers, storage=storage_object) + # Fit elements + with datagrabber_object: + if elements is not None: + for t_element in elements: + mc.fit(datagrabber_object[t_element]) + else: + for t_element in datagrabber_object: + mc.fit(datagrabber_object[t_element]) + + +def collect(storage: Dict) -> None: + """Collect and store data. + + Parameters + ---------- + storage : dict + Storage to use. Must have a key "kind" with the kind of + storage to use. All other keys are passed to the storage + init function. + + """ + storage_params = storage.copy() + storage_kind = storage_params.pop("kind") + logger.info(f"Collecting data using {storage_kind}") + logger.debug(f"\tStorage params: {storage_params}") + storage_object = build( + step="storage", + name=storage_kind, + baseclass=BaseFeatureStorage, + init_params=storage_params, + ) + storage_object = typing.cast(BaseFeatureStorage, storage_object) + logger.debug("Running storage.collect()") + storage_object.collect() + logger.info("Collect done") + + +def queue( + config: Dict, + kind: str, + jobname: str = "junifer_job", + overwrite: bool = False, + elements: Union[str, List[Union[str, Tuple]], Tuple, None] = None, + **kwargs: Union[str, int, bool], +) -> None: # pragma : no cover + """Queue a job to be executed later. + + Parameters + ---------- + config : dict + The configuration to be used for queueing the job. + kind : {"HTCondor", "SLURM"} + The kind of job queue system to use. + jobname : str, optional + The name of the job (default "junifer_job"). + overwrite : bool, optional + Whether to overwrite if job directory already exists (default False). + elements : str or tuple or list of str or tuple, optional + Element(s) to process. Will be used to index the datagrabber + (default None). + **kwargs : dict + The keyword arguments to pass to the job queue system. + + Raises + ------ + ValueError + If the value of `kind` is invalid. + + """ + # Create a folder within the CWD to store the job files / config + cwd = Path.cwd() + jobdir = cwd / "junifer_jobs" / jobname + logger.info(f"Creating job in {str(jobdir.absolute())}") + if jobdir.exists(): + if overwrite is not True: + raise_error( + f"Job folder for {jobname} already exists. " + "This error is raised to prevent overwriting job files " + "that might be scheduled but not yet executed. " + f"Either delete the directory {str(jobdir.absolute())} " + "or set overwrite=True." + ) + else: + logger.info( + f"Deleting existing job directory at {str(jobdir.absolute())}" + ) + shutil.rmtree(jobdir) + jobdir.mkdir(exist_ok=True, parents=True) + + yaml_config = jobdir / "config.yaml" + logger.info(f"Writing YAML config to {str(yaml_config.absolute())}") + with open(yaml_config, "w") as f: + f.write(yaml.dump(config)) + + # Get list of elements + if elements is None: + if "elements" in config: + elements = config["elements"] + else: + # If no elements are specified, use all elements from the + # datagrabber + datagrabber = _get_datagrabber(config["datagrabber"]) + with datagrabber as dg: + elements = dg.get_elements() + + # TODO: Fix typing of elements + if not isinstance(elements, List): + elements = [elements] # type: ignore + + typing.cast(List[Union[str, Tuple]], elements) + + if kind == "HTCondor": + _queue_condor( + jobname=jobname, + jobdir=jobdir, + yaml_config=yaml_config, + elements=elements, # type: ignore + **kwargs, + ) + elif kind == "SLURM": + _queue_slurm( + jobname=jobname, + jobdir=jobdir, + yaml_config=yaml_config, + elements=elements, # type: ignore + **kwargs, + ) + else: + raise ValueError(f"Unknown queue kind: {kind}") + + logger.info("Queue done") + + +def _queue_condor( + jobname: str, + jobdir: Path, + yaml_config: Path, + elements: List[Union[str, Tuple]], + env: Optional[Dict[str, str]] = None, + mem: str = "8G", + cpus: int = 1, + disk: str = "1G", + extra_preamble: str = "", + verbose: str = "info", + collect: bool = True, + submit: bool = False, +) -> None: + """Submit job to HTCondor. + + Parameters + ---------- + jobname : str + The name of the job. + jobdir : pathlib.Path + The path to the job directory. + yaml_config : pathlib.Path + The path to the YAML config file. + elements : list of str or tuple + Element(s) to process. Will be used to index the datagrabber. + env : dict, optional + The environment variables passed as dictionary (default None). + mem : str, optional + The size of memory (RAM) to use (default "8G"). + cpus : int, optional + The number of CPU cores to use (default 1). + disk : str, optional + The size of disk (HDD or SSD) to use (default "1G"). + extra_preamble : str, optional + Extra commands to pass to HTCondor (default ""). + verbose : str, optional + The level of verbosity (default "info"). + collect : bool, optional + Whether to submit "collect" task for junifer (default True). + submit : bool, optional + Whether to submit the jobs. In any case, .dag files will be created + for submission (default False). + + Raises + ------ + ValueError + If the value of `env` is invalid. + + """ + logger.debug("Creating HTCondor job") + run_junifer_args = ( + f"run {str(yaml_config.absolute())} " + f"--verbose {verbose} --element $(element)" + ) + collect_junifer_args = ( + f"collect {str(yaml_config.absolute())} --verbose {verbose} " + ) + + # Set up the env_name, executable and arguments according to the + # environment type + if env is None: + env = {"kind": "local"} + if env["kind"] == "conda": + env_name = env["name"] + executable = "run_conda.sh" + arguments = f"{env_name} junifer" + # TODO: Copy run_conda.sh to jobdir + exec_path = jobdir / executable + shutil.copy(Path(__file__).parent / "res" / executable, exec_path) + make_executable(exec_path) + elif env["kind"] == "venv": + env_name = env["name"] + executable = "run_venv.sh" + arguments = f"{env_name} junifer" + # TODO: Copy run_venv.sh to jobdir + elif env["kind"] == "local": + executable = "junifer" + arguments = "" + else: + raise ValueError(f'Unknown env kind: {env["kind"]}') + + # Create log directory + log_dir = jobdir / "logs" + log_dir.mkdir(exist_ok=True, parents=True) + + # Add preamble data + run_preamble = f""" + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = {cpus} + request_memory = {mem} + request_disk = {disk} + + # Executable + initial_dir = {str(jobdir.absolute())} + executable = $(initial_dir)/{executable} + transfer_executable = False + + arguments = {arguments} {run_junifer_args} + + {extra_preamble} + + # Logs + log = {str(log_dir.absolute())}/junifer_run_$(element).log + output = {str(log_dir.absolute())}/junifer_run_$(element).out + error = {str(log_dir.absolute())}/junifer_run_$(element).err + """ + + submit_run_fname = jobdir / f"run_{jobname}.submit" + submit_collect_fname = jobdir / f"collect_{jobname}.submit" + dag_fname = jobdir / f"{jobname}.dag" + + # Write to run submit files + with open(submit_run_fname, "w") as submit_file: + submit_file.write(run_preamble) + submit_file.write("queue\n") + + collect_preamble = f""" + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = {cpus} + request_memory = {mem} + request_disk = {disk} + + # Executable + initial_dir = {str(jobdir.absolute())} + executable = $(initial_dir)/{executable} + transfer_executable = False + + arguments = {arguments} {collect_junifer_args} + + {extra_preamble} + + # Logs + log = {str(log_dir.absolute())}/junifer_collect.log + output = {str(log_dir.absolute())}/junifer_collect.out + error = {str(log_dir.absolute())}/junifer_collect.err + """ + + # Now create the collect submit file + with open(submit_collect_fname, "w") as submit_file: + submit_file.write(collect_preamble) # Eval preamble here + submit_file.write("queue\n") + + with open(dag_fname, "w") as dag_file: + # Get all subject and session names from file list + for i_job, t_elem in enumerate(elements): + dag_file.write(f"JOB run{i_job} {submit_run_fname}\n") + dag_file.write(f'VARS run{i_job} element="{t_elem}"\n\n') + if collect is True: + dag_file.write(f"JOB collect {submit_collect_fname}\n") + dag_file.write("PARENT ") + for i_job, _t_elem in enumerate(elements): + dag_file.write(f"run{i_job} ") + dag_file.write("CHILD collect\n\n") + + # Submit job(s) + if submit is True: + logger.info("Submitting HTCondor job") + subprocess.run(["condor_submit_dag", dag_fname]) + logger.info("HTCondor job submitted") + else: + cmd = f"condor_submit_dag {str(dag_fname.absolute())}" + logger.info( + f"HTCondor job files created, to submit the job, run `{cmd}`" + ) + + +def _queue_slurm( + jobname: str, + jobdir: Path, + yaml_config: Path, + elements: List[Union[str, Tuple]], +) -> None: + """Submit job to SLURM. + + Parameters + ---------- + jobname : str + The name of the job. + jobdir : pathlib.Path + The path to the job directory. + yaml_config : pathlib.Path + The path to the YAML config file. + elements : str or tuple or list[str or tuple], optional + Element(s) to process. Will be used to index the datagrabber + (default None). + + """ + pass + # logger.debug("Creating SLURM job") + # run_junifer_args = ( + # f"run {str(yaml_config.absolute())} " + # f"--verbose {verbose} --element $(element)" + # ) + # collect_junifer_args = \ + # f"collect {str(yaml_config.absolute())} --verbose {verbose} " + + # # Set up the env_name, executable and arguments according to the + # # environment type + # if env is None: + # env = { + # "kind": "local", + # } + # if env["kind"] == "conda": + # env_name = env["name"] + # executable = "run_conda.sh" + # arguments = f"{env_name} junifer" + # # TODO: Copy run_conda.sh to jobdir + # exec_path = jobdir / executable + # shutil.copy(Path(__file__).parent / "res" / executable, exec_path) + # make_executable(exec_path) + # elif env["kind"] == "venv": + # env_name = env["name"] + # executable = "run_venv.sh" + # arguments = f"{env_name} junifer" + # # TODO: Copy run_venv.sh to jobdir + # elif env["kind"] == "local": + # executable = "junifer" + # arguments = "" + # else: + # raise ValueError(f"Unknown env kind: {env['kind']}") + + # # Create log directory + # log_dir = jobdir / 'logs' + # log_dir.mkdir(exist_ok=True, parents=True) + + # # Add preamble data + # run_preamble = f""" + # #!/bin/bash + + # #SBATCH --job-name={} + # #SBATCH --account={} + # #SBATCH --partition={} + # #SBATCH --time={} + # #SBATCH --ntasks={} + # #SBATCH --cpus-per-task={cpus} + # #SBATCH --mem-per-cpu={mem} + # #SBATCH --mail-type={} + # #SBATCH --mail-user={} + # #SBATCH --output={} + # #SBATCH --error={} + + # # Executable + # initial_dir = {str(jobdir.absolute())} + # executable = $(initial_dir)/{executable} + # transfer_executable = False + + # arguments = {arguments} {run_junifer_args} + + # {extra_preamble} + + # # Logs + # log = {str(log_dir.absolute())}/junifer_run_$(element).log + # output = {str(log_dir.absolute())}/junifer_run_$(element).out + # error = {str(log_dir.absolute())}/junifer_run_$(element).err + # """ + + # submit_run_fname = jobdir / f'run_{jobname}.sh' + # submit_collect_fname = jobdir / f'collect_{jobname}.sh' + + # # Write to run submit files + # with open(submit_run_fname, 'w') as submit_file: + # submit_file.write(run_preamble) + # submit_file.write('queue\n') + + # collect_preamble = f""" + # # The environment + # universe = vanilla + # getenv = True + + # # Resources + # request_cpus = {cpus} + # request_memory = {mem} + # request_disk = {disk} + + # # Executable + # initial_dir = {str(jobdir.absolute())} + # executable = $(initial_dir)/{executable} + # transfer_executable = False + + # arguments = {arguments} {collect_junifer_args} + + # {extra_preamble} + + # # Logs + # log = {str(log_dir.absolute())}/junifer_collect.log + # output = {str(log_dir.absolute())}/junifer_collect.out + # error = {str(log_dir.absolute())}/junifer_collect.err + # """ + + # # Now create the collect submit file + # with open(submit_collect_fname, 'w') as submit_file: + # submit_file.write(collect_preamble) # Eval preamble here + # submit_file.write('queue\n') + + # with open(dag_fname, 'w') as dag_file: + # # Get all subject and session names from file list + # for i_job, t_elem in enumerate(elements): + # dag_file.write(f'JOB run{i_job} {submit_run_fname}\n') + # dag_file.write(f'VARS run{i_job} element="{t_elem}"\n\n') + # if collect is True: + # dag_file.write(f'JOB collect {submit_collect_fname}\n') + # dag_file.write('PARENT ') + # for i_job, _t_elem in enumerate(elements): + # dag_file.write(f'run{i_job} ') + # dag_file.write('CHILD collect\n\n') + + # # Submit job(s) + # if submit is True: + # logger.info('Submitting SLURM job') + # subprocess.run(['condor_submit_dag', dag_fname]) + # logger.info('HTCondor SLURM submitted') + # else: + # cmd = f"condor_submit_dag {str(dag_fname.absolute())}" + # logger.info( + # f"SLURM job files created, to submit the job, run `{cmd}`" + # ) diff --git a/junifer/api/parser.py b/junifer/api/parser.py new file mode 100644 index 000000000..470f39bed --- /dev/null +++ b/junifer/api/parser.py @@ -0,0 +1,51 @@ +"""Provide functions for parser.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import importlib +from pathlib import Path +from typing import Dict, Union + +import yaml + +from ..utils.logging import logger, raise_error + + +def parse_yaml(filepath: Union[str, Path]) -> Dict: + """Parse YAML. + + Parameters + ---------- + filepath : str or pathlib.Path + The filepath to read from. + + Returns + ------- + dict + The contents represented as dictionary. + + """ + # Convert str to Path + if not isinstance(filepath, Path): + filepath = Path(filepath) + + logger.info(f"Parsing yaml file: {str(filepath.absolute())}") + # Filepath existence check + if not filepath.exists(): + raise_error(f"File does not exist: {str(filepath.absolute())}") + # Filepath reading + with open(filepath, "r") as f: + contents = yaml.safe_load(f) + # Autload modules + if "with" in contents: + to_load = contents["with"] + # Convert autload modules to list + if not isinstance(to_load, list): + to_load = [to_load] + for t_module in to_load: + logger.info(f"Importing module: {t_module}") + importlib.import_module(t_module) + + return contents diff --git a/junifer/api/pipeline.py b/junifer/api/pipeline.py deleted file mode 100644 index 0bf3a5c23..000000000 --- a/junifer/api/pipeline.py +++ /dev/null @@ -1,60 +0,0 @@ -# Authors: Federico Raimondo -# Leonard Sasse -# License: AGPL -from ..utils.logging import raise_error, logger - -_valid_steps = [ - 'datagrabber', 'datareader', 'preprocessing', 'marker', 'storage'] - -_registry = {x: {} for x in _valid_steps} - - -def register(step, name, klass): - """Register a function to be used in a pipeline step - - Parameters - ---------- - step : str - Name of the step - name : str - Name of the function - klass : class - Class to be registered - """ - if step not in _valid_steps: - raise_error(f'Invalid step: {step}', ValueError) - logger.info(f'Registering {name} in {step}') - _registry[step][name] = klass - - -def run_pipeline( - workdir, datagrabber, element, markers, storage, source_params=None, - storage_params=None): - """Run the pipeline on the selected element - - Parameters - ---------- - workdir : str or path-like object - Directory where the pipeline will be executed - datagrabber : str - Name of the datagrabber to use - element : str - Name of the element to process. Will be used to index the datagrabber. - markers : list of dict - List of markers to extract. Each marker is a dict with at least two - keys: 'name' and 'kind'. The 'name' key is used to name the output - marker. The 'kind' key is used to specify the kind of marker to - extract. The rest of the keys are used to pass parameters to the - marker calculation. - storage: str - Name of the storage to use. - source_params : dict - Parameters to pass to the datagrabber. - storage_params: dict - Parameters to pass to the storage. - """ - if source_params is None: - source_params = {} - - if storage_params is None: - storage_params = {} diff --git a/junifer/api/registry.py b/junifer/api/registry.py new file mode 100644 index 000000000..c1398d211 --- /dev/null +++ b/junifer/api/registry.py @@ -0,0 +1,146 @@ +"""Provide functions for registry.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from typing import TYPE_CHECKING, Dict, List, Optional, Union + +from ..utils.logging import logger, raise_error + + +if TYPE_CHECKING: + from ..datagrabber.base import BaseDataGrabber + from ..pipeline import PipelineStepMixin + from ..storage.base import BaseFeatureStorage + +# Define valid steps for operation +_valid_steps = [ + "datagrabber", + "datareader", + "preprocessing", + "marker", + "storage", +] + +# Define registry for valid steps +_registry = {x: {} for x in _valid_steps} + + +def register(step: str, name: str, klass: type) -> None: + """Register a function to be used in a pipeline step. + + Parameters + ---------- + step : str + Name of the step. + name : str + Name of the function. + klass : class + Class to be registered. + + """ + # Verify step + if step not in _valid_steps: + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + + logger.info(f"Registering {name} in {step}") + _registry[step][name] = klass + + +def get_step_names(step: str) -> List: + """Get the names of the registered functions for a given step. + + Parameters + ---------- + step : str + Name of the step. + + Returns + ------- + list + List of registered function names. + + """ + # Verify step + if step not in _valid_steps: + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + + return list(_registry[step].keys()) + + +def get_class(step: str, name: str) -> type: + """Get the class of the registered function for a given step. + + Parameters + ---------- + step : str + Name of the step. + name : str + Name of the function. + + Returns + ------- + class + Registered function class. + + """ + # Verify step + if step not in _valid_steps: + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + # Verify step name + if name not in _registry[step]: + raise_error(msg=f"Invalid name: {name}", klass=ValueError) + + return _registry[step][name] + + +def build( + step: str, + name: str, + baseclass: type, + init_params: Optional[Dict] = None, +) -> Union["BaseDataGrabber", "PipelineStepMixin", "BaseFeatureStorage"]: + """Ensure that the given object is an instance of the given class. + + Parameters + ---------- + step : str + Name of the step. + name : str + Name of the function. + baseclass : class + Class to be checked against. + init_parms : dict, optional + Parameters to pass to the base class constructor (default None). + + Returns + ------- + object + An instance of the given base class. + + Raises + ------ + ValueError + If the created object with the given name is not an instance of the + base class. + + """ + # Set default init parameters + if init_params is None: + init_params = {} + # Get class of the registered function + klass = get_class(step=step, name=name) + # Create instance of the class + object_ = klass(**init_params) + # Verify created instance belongs to the base class + if not isinstance(object_, baseclass): + raise_error( + msg=( + f"Invalid {step} ({object_.__class__.__name__}). " + f"Must inherit from {baseclass.__name__}" + ), + klass=ValueError, + ) + return object_ diff --git a/junifer/api/res/run_conda.sh b/junifer/api/res/run_conda.sh new file mode 100755 index 000000000..f07c86c6f --- /dev/null +++ b/junifer/api/res/run_conda.sh @@ -0,0 +1,17 @@ +#!/bin/bash + +if [ $# -lt 2 ]; then + echo "This script is meant to run a command within a python environment" + echo "It needs at least 2 parameters." + echo "The first one must be the environment name." + echo "The rest will be the command" + exit 255 +fi + +eval "$(conda shell.bash hook)" +env_name=$1 +echo "Activating ${env_name}" +conda activate "$1" +shift 1 +echo "Running ${*} in virtual environment" +"$@" \ No newline at end of file diff --git a/junifer/api/tests/data/gmd_mean.yaml b/junifer/api/tests/data/gmd_mean.yaml new file mode 100644 index 000000000..c492cc559 --- /dev/null +++ b/junifer/api/tests/data/gmd_mean.yaml @@ -0,0 +1,15 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +elements: [1, 2] +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db + diff --git a/junifer/api/tests/data/gmd_mean_htcondor.yaml b/junifer/api/tests/data/gmd_mean_htcondor.yaml new file mode 100644 index 000000000..fddf7aa9b --- /dev/null +++ b/junifer/api/tests/data/gmd_mean_htcondor.yaml @@ -0,0 +1,20 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db +queue: + jobname: TestHTCondorQueue + kind: HTCondor + env: + kind: conda + name: junifer + mem: 8G \ No newline at end of file diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py new file mode 100644 index 000000000..380196f12 --- /dev/null +++ b/junifer/api/tests/test_cli.py @@ -0,0 +1,70 @@ +"""Provide tests for cli.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import Tuple + +import pytest +import yaml +from click.testing import CliRunner + +from junifer.api.cli import collect, run + + +# Create click test runner +runner = CliRunner() + + +# TODO: adapt elements to take arrays +@pytest.mark.parametrize( + "elements", + [ + ("sub-01", "sub-02", "sub-03"), + ("sub-01", "sub-02", "sub-04"), + ], +) +def test_run_and_collect_commands( + tmp_path: Path, elements: Tuple[str, ...] +) -> None: + """Test run and collect commands.""" + # Get test config + infile = Path(__file__).parent / "data" / "gmd_mean.yaml" + # Read test config + with open(infile, mode="r") as f: + contents = yaml.safe_load(f) + # Working directory + workdir = tmp_path / "workdir" + contents["workdir"] = str(workdir.absolute()) + # Output directory + outdir = tmp_path / "outdir" + # Storage + contents["storage"]["uri"] = str(outdir.absolute()) + # Write new test config + outfile = tmp_path / "in.yaml" + with open(outfile, mode="w") as f: + yaml.dump(contents, f) + # Run command arguments + run_args = [ + str(outfile.absolute()), + "--verbose", + "debug", + "--element", + elements[0], + "--element", + elements[1], + "--element", + elements[2], + ] + # Invoke run command + run_result = runner.invoke(run, run_args) + # Check + assert run_result.exit_code == 0 + # Collect command arguments + collect_args = [str(outfile.absolute()), "--verbose", "debug"] + # Invoke collect command + collect_result = runner.invoke(collect, collect_args) + # Check + assert collect_result.exit_code == 0 diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py new file mode 100644 index 000000000..3b175c901 --- /dev/null +++ b/junifer/api/tests/test_functions.py @@ -0,0 +1,157 @@ +"""Provide tests for functions.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest + +import junifer.testing.registry # noqa: F401 +from junifer.api.functions import collect, run +from junifer.api.registry import build +from junifer.datagrabber.base import BaseDataGrabber + + +# Define datagrabber +datagrabber = { + "kind": "OasisVBMTestingDatagrabber", +} + +# Define markers +markers = [ + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, +] + +# Define storage +storage = { + "kind": "SQLiteFeatureStorage", +} + + +def test_run_single_element(tmp_path: Path) -> None: + """Test run function with single element. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Create working directory + workdir = tmp_path / "workdir_single" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri # type: ignore + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=["sub-01"], + ) + # Check files + files = list(outdir.glob("*.db")) + assert len(files) == 1 + + +def test_run_multi_element(tmp_path: Path) -> None: + """Test run function with multi element. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Create working directory + workdir = tmp_path / "workdir_multi" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri # type: ignore + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=["sub-01", "sub-03"], + ) + # Check files + files = list(outdir.glob("*.db")) + assert len(files) == 2 + + +def test_run_and_collect(tmp_path: Path) -> None: + """Test run and collect functions. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Create working directory + workdir = tmp_path / "workdir" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri # type: ignore + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + ) + # Get datagrabber + dg = build( + step="datagrabber", name=datagrabber["kind"], baseclass=BaseDataGrabber + ) + elements = dg.get_elements() # type: ignore + # This should create 10 files + files = list(outdir.glob("*.db")) + assert len(files) == len(elements) + # But the test.db file should not exist + assert not uri.exists() + # Collect in storage + collect(storage) + # Now the file exists + assert uri.exists() + + +@pytest.mark.skip(reason="HTCondor not installed on system.") +def test_queue_condor() -> None: + """Test job queueing in HTCondor.""" + pass + + +@pytest.mark.skip(reason="SLURM not installed on system.") +def test_queue_slurm() -> None: + """Test job queueing in SLURM.""" + pass diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py new file mode 100644 index 000000000..e61ff9fdb --- /dev/null +++ b/junifer/api/tests/test_parser.py @@ -0,0 +1,77 @@ +"""Provide tests for parser.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import sys +from pathlib import Path + +import pytest + +from junifer.api.parser import parse_yaml + + +def test_parse_yaml_failure() -> None: + """Test YAML parsing failure.""" + with pytest.raises(ValueError, match="does not exist"): + parse_yaml("foo.yaml") + + +def test_parse_yaml_success(tmp_path: Path) -> None: + """Test YAML parsing success. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Write test file + fname = tmp_path / "test_parse_yaml_success.yaml" + fname.write_text("foo: bar") + # Check test file + contents = parse_yaml(fname) + assert "foo" in contents + assert contents["foo"] == "bar" + + +def test_parse_yaml_success_with_module_autoload(tmp_path: Path) -> None: + """Test YAML parsing with single module autoload success. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Write test file + fname = tmp_path / "test_parse_yaml_with_single_module_autoload.yaml" + fname.write_text("foo: bar\nwith: numpy") + # Check test file + contents = parse_yaml(fname) + assert "foo" in contents + assert contents["foo"] == "bar" + assert "with" in contents + assert contents["with"] == "numpy" + assert "numpy" in sys.modules + assert "junifer.configs.wrong_config" not in sys.modules + + +def test_parse_yaml_failure_with_multi_module_autoload(tmp_path: Path) -> None: + """Test YAML parsing with multi module autoload failure. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Write test file + fname = tmp_path / "test_parse_yaml_with_multi_module_autoload.yaml" + fname.write_text( + "foo: bar\nwith:\n - numpy\n - junifer.testing.wrong_config" + ) + # Check test file + with pytest.raises(ImportError, match="wrong_config"): + parse_yaml(fname) diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py new file mode 100644 index 000000000..7a45c0d5f --- /dev/null +++ b/junifer/api/tests/test_registry.py @@ -0,0 +1,136 @@ +"""Provide tests for registry.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL +import logging +from abc import ABC +from typing import Type + +import pytest + +from junifer.api.registry import build, get_class, get_step_names, register +from junifer.datagrabber import PatternDataGrabber +from junifer.storage import SQLiteFeatureStorage + + +def test_register_invalid_step(): + """Test register invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + register(step="foo", name="bar", klass=str) + + +# TODO: improve parametrization +@pytest.mark.parametrize( + "step, name, klass", + [ + ("datagrabber", "pattern-dg", PatternDataGrabber), + ("storage", "sqlite-storage", SQLiteFeatureStorage), + ], +) +def test_register( + caplog: pytest.LogCaptureFixture, step: str, name: str, klass: Type +) -> None: + """Test register. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + step : str + The parametrized name of the step. + name : str + The parametrized name of the function. + klass : str + The parametrized name of the base class. + + """ + with caplog.at_level(logging.INFO): + # Register + register(step=step, name=name, klass=klass) + # Check logging message + assert "Registering" in caplog.text + + +def test_get_step_names_invalid_step() -> None: + """Test get step name invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + get_step_names(step="foo") + + +def test_get_step_names_absent() -> None: + """Test get step names for absent name.""" + # Get step names for datagrabber + datagrabbers = get_step_names(step="datagrabber") + # Check for datagrabber step name + assert "bar" not in datagrabbers + + +def test_get_step_names() -> None: + """Test get step names.""" + # Register datagrabber + register(step="datagrabber", name="bar", klass=str) + # Get step names for datagrabber + datagrabbers = get_step_names(step="datagrabber") + # Check for datagrabber step name + assert "bar" in datagrabbers + + +def test_get_class_invalid_step() -> None: + """Test get class invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + get_class(step="foo", name="bar") + + +def test_get_class_invalid_name() -> None: + """Test get class invalid function name.""" + with pytest.raises(ValueError, match="Invalid name:"): + get_class(step="datagrabber", name="foo") + + +# TODO: enable parametrization +def test_get_class(): + """Test get class.""" + # Register datagrabber + register(step="datagrabber", name="bar", klass=str) + # Get class + obj = get_class(step="datagrabber", name="bar") + assert obj == str + + +# TODO: possible parametrization? +def test_build(): + """Test building objects from names.""" + import numpy as np + + # Define abstract base class + class SuperClass(ABC): + pass + + # Define concrete class + class ConcreteClass(SuperClass): + def __init__(self, value=1): + self.value = value + + # Register + register(step="datagrabber", name="concrete", klass=ConcreteClass) + + # Build + obj = build(step="datagrabber", name="concrete", baseclass=SuperClass) + assert isinstance(obj, ConcreteClass) + assert obj.value == 1 + + # Build + obj = build( + step="datagrabber", + name="concrete", + baseclass=SuperClass, + init_params={"value": 2}, + ) + assert isinstance(obj, ConcreteClass) + assert obj.value == 2 + + # Check error + with pytest.raises(ValueError, match="Must inherit"): + build(step="datagrabber", name="concrete", baseclass=np.ndarray) diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index d54e1ee68..fa16361d3 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,65 +1,44 @@ -from ..datagrabber import DataladDataGrabber +"""Provide class for juseless datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import Union + from ..api.decorators import register_datagrabber +from ..datagrabber import PatternDataladDataGrabber @register_datagrabber -class JuselessUKBVBM(DataladDataGrabber): - """Juseless UKB VMG DataGrabber class. +class JuselessDataladUKBVBM(PatternDataladDataGrabber): + """Juseless UKB VBM DataGrabber class. Implements a DataGrabber to access the UKB VBM data in Juseless. + Parameters + ----------- + datadir : str or pathlib.Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + """ - def __init__(self, datadir=None): - """Initialize a JuselessUKBVBM object. - - Parameters - ---------- - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - """ - uri = 'ria+http://ukb.ds.inm7.de#~cat_m0wp1' - rootdir = 'm0wp1' - types = ['VBM_GM'] + def __init__(self, datadir: Union[str, Path, None] = None) -> None: + """Initialize the class.""" + uri = "ria+http://ukb.ds.inm7.de#~cat_m0wp1" + rootdir = "m0wp1" + types = ["VBM_GM"] + replacements = ["subject", "session"] + patterns = {"VBM_GM": "m0wp1sub-{subject}_ses-{session}_T1w.nii.gz"} super().__init__( - types=types, datadir=datadir, uri=uri, rootdir=rootdir) - - def get_elements(self): - """Get the list of subjects in the dataset. - - Returns - ------- - elements : list[str] - The list of subjects in the dataset. - """ - elems = [] - for x in self.datadir.glob('*._T1w.nii.gz'): - sub, ses = x.name.split('_') - sub = sub.replace('m0wp1', '') - ses = ses[:5] - elems.append((sub, ses)) - return elems - - def __getitem__(self, element): - """Index one element in the dataset. - - Parameters - ---------- - element : tuple[str, str] - The element to be indexed. First element in the tuple is the - subject, second element is the session. - - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - sub, ses = element - out = {} - - out['VBM_GM'] = self.datadir / f'm0wp1{sub}_{ses}_T1w.nii.gz' - self._dataset_get(out) - return out + types=types, + datadir=datadir, + uri=uri, + rootdir=rootdir, + replacements=replacements, + patterns=patterns, + ) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 7c5131327..e451c01c2 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,15 +1,47 @@ +"""Provide tests for juseless datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + import socket + import pytest -from junifer.configs.juseless import JuselessUKBVBM - -if socket.gethostname() != 'juseless': - pytest.skip('This tests are only for juseless', allow_module_level=True) +from junifer.configs.juseless import JuselessDataladUKBVBM +from junifer.datagrabber.hcp import DataladHCP1200 +from junifer.utils.logging import configure_logging -def test_juselessukbvbm_datagrabber(): - with JuselessUKBVBM() as dg: - out = dg[('sub-2670511', 'ses-2')] - assert 'VBM_GM' in out - assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz' - assert out['VBM_GM'].exists() +# Check if the test is running on juseless +if socket.gethostname() != "juseless": + pytest.skip("These tests are only for juseless", allow_module_level=True) + +configure_logging(level="DEBUG") + + +def test_juselessdataladukbvbm_datagrabber() -> None: + """Test datalad UKBVBM datagrabber.""" + with JuselessDataladUKBVBM() as dg: + all_elements = dg.get_elements() + test_element = all_elements[0] + out = dg[test_element] + assert "VBM_GM" in out + assert ( + out["VBM_GM"]["path"].name + == f"m0wp1sub-{test_element[0]}_ses-{test_element[1]}_T1w.nii.gz" + ) + assert out["VBM_GM"]["path"].exists() + + +def test_juselessdataladhcp_datagrabber() -> None: + """Test datalad HCP datagrabber.""" + with DataladHCP1200() as dg: + all_elements = dg.get_elements() + test_element = all_elements[0] + + out = dg[test_element] + + assert out["BOLD"]["path"].exists() + assert out["BOLD"]["path"].isfile() diff --git a/junifer/data/__init__.py b/junifer/data/__init__.py index 4e8ab386c..79e23f86f 100644 --- a/junifer/data/__init__.py +++ b/junifer/data/__init__.py @@ -1 +1,7 @@ -from .atlases import list_atlases, register_atlas, load_atlas \ No newline at end of file +"""Provide imports for data sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .atlases import list_atlases, register_atlas, load_atlas diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 64e999e77..9b9bfc3b3 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -1,17 +1,30 @@ +"""Provide functions for atlases.""" + # Authors: Federico Raimondo # Vera Komeyer +# Synchon Mandal # License: AGPL -from pathlib import Path + import io -import requests -import numpy as np -import pandas as pd +import shutil +import tempfile +import zipfile +from pathlib import Path +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import nibabel as nib +import numpy as np +import pandas as pd +import requests from nilearn import datasets from ..utils.logging import logger, raise_error + +if TYPE_CHECKING: + from nibabel import Nifti1Image + + """ A dictionary containing all supported atlases and their respective valid parameters. @@ -23,150 +36,198 @@ Optional keys: * 'valid_resolutions': a list of valid resolutions for the atlas (e.g. [1, 2]) """ -_available_atlases = { - 'SUITxSUIT': { - 'family': 'SUIT', - 'sace': 'SUIT' - }, - 'SUITxMNI': { - 'family': 'SUIT', - 'space': 'MNI' - }, - +# TODO: have separate dictionary for built-in +_available_atlases: Dict[str, Dict[Any, Any]] = { + "SUITxSUIT": {"family": "SUIT", "space": "SUIT"}, + "SUITxMNI": {"family": "SUIT", "space": "MNI"}, } +# Add Schaefer atlas info for n_rois in range(100, 1001, 100): for t_net in [7, 17]: - t_name = f'Schaefer{n_rois}x{t_net}' + t_name = f"Schaefer{n_rois}x{t_net}" _available_atlases[t_name] = { - 'family': 'Schaefer', - 'n_rois': n_rois, - 'yeo_networks': t_net, - 'valid_resolutions': [1, 2] + "family": "Schaefer", + "n_rois": n_rois, + "yeo_networks": t_net, } +# Add Tian atlas info +for scale in range(1, 5): + t_name = f"TianxS{scale}x7TxMNI6thgeneration" + _available_atlases[t_name] = { + "family": "Tian", + "scale": scale, + "magneticfield": "7T", + "space": "MNI6thgeneration", + } + t_name = f"TianxS{scale}x3TxMNI6thgeneration" + _available_atlases[t_name] = { + "family": "Tian", + "scale": scale, + "magneticfield": "3T", + "space": "MNI6thgeneration", + } + t_name = f"TianxS{scale}x3TxMNInonlinear2009cAsym" + _available_atlases[t_name] = { + "family": "Tian", + "scale": scale, + "magneticfield": "3T", + "space": "MNInonlinear2009cAsym", + } -def register_atlas(name, atlas_path, atl_labels, overwrite=False): +def register_atlas( + name: str, + atlas_path: Union[str, Path], + atl_labels: List[str], + overwrite: bool = False, +) -> None: """Register a custom user atlas. Parameters ---------- name : str The name of the atlas. - atlas_path : str + atlas_path : str or pathlib.Path The path to the atlas file. - atl_labels : list(str) + atl_labels : list of str The list of labels for the atlas. - overwrite : bool - If True, overwrite an existing atlas with the same name. Defaults to - False. + overwrite : bool, optional + If True, overwrite an existing atlas with the same name. + Does not apply to built-in atlases (default False). Raises ------ ValueError If the atlas name is already registered and overwrite is set to False or if the atlas name is a built-in atlas. + """ + # Check for attempt of overwriting built-in atlases if name in _available_atlases: if overwrite is True: - logger.info(f'Overwritting {name} atlas') - if _available_atlases[name]['family'] != 'CustomUserAtlas': + logger.info(f"Overwriting {name} atlas") + if _available_atlases[name]["family"] != "CustomUserAtlas": raise_error( - f'Cannot overwrite {name} atlas. It is a built-in atlas.') + f"Cannot overwrite {name} atlas. It is a built-in atlas." + ) else: raise_error( - f'Atlas {name} already registered. Set `overwrite=True` to ' - 'update its value.') + f"Atlas {name} already registered. Set `overwrite=True` to " + "update its value." + ) + # Convert str to Path + if not isinstance(atlas_path, Path): + atlas_path = Path(atlas_path) + # Add user atlas info _available_atlases[name] = { - 'path': atlas_path, 'labels': atl_labels, - 'family': 'CustomUserAtlas'} + "path": str(atlas_path.absolute()), + "labels": atl_labels, + "family": "CustomUserAtlas", + } -def list_atlases(): - """ - List all the available atlases. +def list_atlases() -> List[str]: + """List all the available atlases. Returns ------- - out : list(str) or dict - A list or dict with all available atlases. + list of str + A list with all available atlases. + """ return sorted(_available_atlases.keys()) -def _check_resolution(resolution, valid_resolution): - if resolution is None: - return None - if resolution not in valid_resolution: - raise ValueError(f'Invalid resolution: {resolution}') - return resolution +# def _check_resolution(resolution, valid_resolution): +# if resolution is None: +# return None +# if resolution not in valid_resolution: +# raise ValueError(f'Invalid resolution: {resolution}') +# return resolution -def load_atlas(name, atlas_dir=None, resolution=None, path_only=False, - **kwargs): - """ - Loads a brain atlas (including a label file). - If it is built-in atlas and file is not present in the `atlas_dir` +# TODO: keyword arguments are not passed, check +def load_atlas( + name: str, + atlas_dir: Union[str, Path, None] = None, + resolution: Optional[float] = None, + path_only: bool = False, +) -> Tuple[Optional["Nifti1Image"], List[str], Path]: + """Load a brain atlas (including a label file). + + If it is a built-in atlas and file is not present in the `atlas_dir` directory, it will be downloaded. Parameters ---------- name : str - The name of the atlas. - Check valid options by calling `list_atlases`. - atlas_dir: path - Path where the atlas files are stored. - Defaults to: $HOME/junifer/data/atlas - resolution : int - The (desired) resolution of the atlas to load. If its not available, + The name of the atlas. Check valid options by calling `list_atlases`. + atlas_dir : str or pathlib.Path, optional + Path where the atlas files are stored. The default location is + "$HOME/junifer/data/atlas" (default None). + resolution : float, optional + The desired resolution of the atlas to load. If it is not available, the closest resolution will be loaded. Preferably, use a resolution - higher than the desired one. Defaults to None (load the highest one). - path_only : bool - If True, the atlas image will not be loaded. + higher than the desired one. By default, will load the highest one + (default None). + path_only : bool, optional + If True, the atlas image will not be loaded (default False). - Parameters (optional, atlas dependent) - -------------------------------------- - Use to specify atlas specific keyword arguments. . + Extra Parameters + ---------------- + Use to specify atlas specific keyword arguments. - Schaefer : - n_rois (required) : int - Granularity of atlas to be used. Valid values: between 100 and 1000 - (included) in steps of 100. - yeo_network (optional) : int - Number of yeo networks to use. Valid values: 7, 17. Defaults to 7. - - Tian : - # TODO add - - SUIT : - space (optional) : str - Space of atlas can be either 'MNI' or 'SUIT' (for more information - see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to - 'MNI'. + - Schaefer : + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000} + Granularity of atlas to be used. + yeo_network : {7, 17}, optional + Number of yeo networks to use (default 7). + - Tian : + scale : {1, 2, 3, 4} + Scale of atlas (defines granularity). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more information + see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + - SUIT : + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). Returns ------- - atlas_img : niimg-like object or None + niimg-like object or None Loaded atlas image. - atlas_labels : List of str + list of str Atlas labels. - atlas_fname : Path + pathlib.Path File path to the atlas image. + """ + # Invalid atlas name + if name not in _available_atlases: + raise_error( + f"Atlas {name} not found. Valid options are: {list_atlases()}" + ) - atlas_definition = _available_atlases[name] - t_family = atlas_definition.pop('family') + atlas_definition = _available_atlases[name].copy() + t_family = atlas_definition.pop("family") - if t_family == 'CustomUserAtlas': - atlas_fname = atlas_definition['path'] - atlas_labels = atlas_definition['labels'] + if t_family == "CustomUserAtlas": + atlas_fname = Path(atlas_definition["path"]) + atlas_labels = atlas_definition["labels"] else: # retrieve atlases by passing arguments on to _retrieve_atlas() atlas_fname, atlas_labels = _retrieve_atlas( - t_family, out_dir=atlas_dir, **kwargs) + family=t_family, + atlas_dir=atlas_dir, + resolution=resolution, + **atlas_definition, + ) - logger.info( - f'Loading atlas {atlas_fname.as_posix()}') # type: ignore + logger.info(f"Loading atlas {str(atlas_fname.absolute())}") atlas_img = None if path_only is False: @@ -175,75 +236,122 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False, return atlas_img, atlas_labels, atlas_fname -def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): - """ - Retrieves a brain atlas object either from nilearn or a specified online - source. Only returns one atlas per call. Call function multiple times for +def _retrieve_atlas( + family: str, + atlas_dir: Union[str, Path, None] = None, + resolution: Optional[float] = None, + **kwargs, +) -> Tuple[Path, List[str]]: + """Retrieve a brain atlas object from nilearn or a specified online source. + + Only returns one atlas per call. Call function multiple times for different parameter specifications. Only retrieves atlas if it is not yet in atlas_dir. - Parameters (required) - --------------------- + Parameters + ---------- family : str - Specify by name of atlas family, e.g. 'Schaefer'. - atlas_dir: path - Path to where to store the retrieved atlas file. - Defaults to: $HOME/junifer/data/atlas - resolution : int - The (desired) resolution of the atlas to load. If its not available, + The name of the atlas family, e.g. 'Schaefer'. + atlas_dir : str or pathlib.Path, optional + Path where the retrieved atlas file is stored. The default location is + "$HOME/junifer/data/atlas" (default None). + resolution : float, optional + The desired resolution of the atlas to load. If it is not available, the closest resolution will be loaded. Preferably, use a resolution - higher than the desired one. Defaults to None (load the highest one). + higher than the desired one. By default, will load the highest one + (default None). - Parameters (optional, atlas dependent) - -------------------------------------- - Use to specify atlas specific keyword arguments + Extra Parameters + ---------------- + **kwargs + Use to specify atlas specific keyword arguments. - Schaefer : - n_rois (required) : int - Granularity of atlas to be used. Valid values: between 100 and 1000 - (included) in steps of 100. - yeo_network (optional) : int - Number of yeo networks to use [7 or 17]. Defaults to 7. - SUIT : - space (optional) : str - Space of atlas can be either 'MNI' or 'SUIT' (for more information - see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to - 'MNI'. + - Schaefer : + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000} + Granularity of atlas to be used. + yeo_network : {7, 17}, optional + Number of yeo networks to use (default 7). + - Tian : + scale : {1, 2, 3, 4} + Scale of atlas (defines granularity). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more + information see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + - SUIT : + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). Returns ------- - atlas_fname : Path + pathlib.Path File path to the atlas image. - atlas_labels : List of str + list of str Atlas labels. + + Raises + ------ + ValueError + If the atlas name is invalid. + """ if atlas_dir is None: - atlas_dir = Path().home() / 'junifer' / 'data' / 'atlas' + atlas_dir = Path().home() / "junifer" / "data" / "atlas" + # Create default junifer data directory if not present atlas_dir.mkdir(exist_ok=True, parents=True) + # Convert str to Path + elif not isinstance(atlas_dir, Path): + atlas_dir = Path(atlas_dir) logger.info(f"Fetching one of {family} atlas.") - # retrieval details per atlas - if family == 'Schaefer': - atlas_fname, atl_labels = \ - _retrieve_schaefer(atlas_dir, **kwargs) - elif family == 'SUIT': - atlas_fname, atl_labels = \ - _retrieve_suit(atlas_dir, **kwargs) + # Retrieval details per atlas + if family == "Schaefer": + atlas_fname, atl_labels = _retrieve_schaefer( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) + elif family == "SUIT": + atlas_fname, atl_labels = _retrieve_suit( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) + elif family == "Tian": + atlas_fname, atl_labels = _retrieve_tian( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) else: - raise_error( - f"The provided atlas name {family} cannot be retrieved. ") + raise_error(f"The provided atlas name {family} cannot be retrieved.") return atlas_fname, atl_labels -def _closest_resolution(resolution, valid_resolution): - closest = None +def _closest_resolution( + resolution: Optional[float], + valid_resolution: Union[List[float], List[int], np.ndarray], +) -> Union[float, int]: + """Find the closest resolution. + + Parameters + ---------- + resolution : float + The given resolution. + valid_resolution : list of float or np.ndarray + The array of valid resolutions. + + Returns + ------- + float + The closest valid resolution. + + """ + # Convert list of int to numpy.ndarray if not isinstance(valid_resolution, np.ndarray): valid_resolution = np.array(valid_resolution) + if resolution is None: - logger.info('Resolution set to None, using highest resolution.') - closest = np.min(valid_resolution) + logger.info("Resolution set to None, using highest resolution.") + closest = np.min(valid_resolution) elif any(x <= resolution for x in valid_resolution): # Case 1: get the highest closest resolution closest = np.max(valid_resolution[valid_resolution <= resolution]) @@ -254,11 +362,46 @@ def _closest_resolution(resolution, valid_resolution): return closest -def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): - logger.info('Atlas parameters:') - logger.info(f'\tn_rois: {n_rois}') - logger.info(f'\tyeo_network: {yeo_network}') - logger.info(f'\tresolution: {resolution}') +def _retrieve_schaefer( + atlas_dir: Path, + resolution: Optional[float] = None, + n_rois: Optional[int] = None, + yeo_networks: int = 7, +) -> Tuple[Path, List[str]]: + """Retrieve Schaefer atlas. + + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : float, optional + The desired resolution of the atlas to load. If it is not available, + the closest resolution will be loaded. Preferably, use a resolution + higher than the desired one. By default, will load the highest one + (default None). Available resolutions for this atlas are 1mm and 2mm. + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000}, optional + Granularity of the atlas to be used (default None). + yeo_networks : {7, 17}, optional + Number of yeo networks to use (default 7). + + Returns + ------- + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `n_rois` or `yeo_networks` or if + there is a problem fetching the atlas. + + """ + logger.info("Atlas parameters:") + logger.info(f"\tn_rois: {n_rois}") + logger.info(f"\tyeo_networks: {yeo_networks}") + logger.info(f"\tresolution: {resolution}") _valid_n_rois = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000] _valid_networks = [7, 17] @@ -266,57 +409,265 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): if n_rois not in _valid_n_rois: raise_error( - f'The parameter `n_rois` ({n_rois}) needs to be one of the ' - f'following: {_valid_n_rois}') - if yeo_network not in _valid_networks: + f"The parameter `n_rois` ({n_rois}) needs to be one of the " + f"following: {_valid_n_rois}" + ) + if yeo_networks not in _valid_networks: raise_error( - f'The parameter `yeo_network` ({yeo_network}) needs to be one of ' - f'the following: {_valid_networks}') + f"The parameter `yeo_networks` ({yeo_networks}) needs to be one " + f"of the following: {_valid_networks}" + ) resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - atlas_fname = atlas_dir / 'schaefer_2018' / ( - f'Schaefer2018_{n_rois}Parcels_{yeo_network}Networks_order_' - f'FSLMNI152_{resolution}mm.nii.gz') - atlas_lname = atlas_dir / 'schaefer_2018' / ( - f'Schaefer2018_{n_rois}Parcels_{yeo_network}Networks_order.txt') + atlas_fname = ( + atlas_dir + / "schaefer_2018" + / ( + f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order_" + f"FSLMNI152_{resolution}mm.nii.gz" + ) + ) + atlas_lname = ( + atlas_dir + / "schaefer_2018" + / (f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt") + ) - # check existance of atlas + # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): logger.info( - 'At least one of the atlas files is missing. ' - 'Fetching using nilearn.') + "At least one of the atlas files is missing. " + "Fetching using nilearn." + ) datasets.fetch_atlas_schaefer_2018( n_rois=n_rois, - yeo_networks=yeo_network, - resolution_mm=resolution, - data_dir=atlas_dir.as_posix()) + yeo_networks=yeo_networks, + resolution_mm=resolution, # type: ignore we know it's 1 or 2 + data_dir=str(atlas_dir.absolute()), + ) - if not (atlas_fname.exists() and atlas_lname.exists()): - raise_error('There was a problem fetching the atlases.') + if not ( + atlas_fname.exists() and atlas_lname.exists() + ): # pragma: no cover + raise_error("There was a problem fetching the atlases.") # Load labels labels = [ - '_'.join(x.split('_')[1:]) - for x in pd.read_csv( - atlas_lname, sep='\t', header=None).iloc[:, 1].to_list() + "_".join(x.split("_")[1:]) + for x in pd.read_csv(atlas_lname, sep="\t", header=None) + .iloc[:, 1] + .to_list() ] return atlas_fname, labels -def _retrieve_suit(out_dir, resolution, space='MNI'): - logger.info('Atlas parameters:') - logger.info(f'\tspace: {space}') +def _retrieve_tian( + atlas_dir: Path, + resolution: Optional[float] = None, + scale: Optional[int] = None, + space: str = "MNI6thgeneration", + magneticfield: str = "3T", +) -> Tuple[Path, List[str]]: + """Retrieve Tian atlas. - _valid_spaces = ['MNI', 'SUIT'] + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : float, optional + The desired resolution of the atlas to load. If it is not available, + the closest resolution will be loaded. Preferably, use a resolution + higher than the desired one. By default, will load the highest one + (default None). Available resolutions for this atlas depend on the + space and magnetic field. + scale : {1, 2, 3, 4}, optional + Scale of atlas (defines granularity) (default None). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more + information see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + + Returns + ------- + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `scale` or `magneticfield` or `space` + or if there is a problem fetching the atlas. + + """ + # show atlas parameters to user + logger.info("Atlas parameters:") + logger.info(f"\tscale: {scale}") + logger.info(f"\tspace: {space}") + logger.info(f"\tmagneticfield: {magneticfield}") + logger.info(f"\tresolution: {resolution}") + # check validity of atlas parameters + _valid_scales = [1, 2, 3, 4] + if scale not in _valid_scales: + raise_error( + f"The parameter `scale` ({scale}) needs to be one of the " + f"following: {_valid_scales}" + ) + + _valid_resolutions = [] # avoid pylance error + if magneticfield == "3T": + _valid_spaces = ["MNI6thgeneration", "MNInonlinear2009cAsym"] + if space == "MNI6thgeneration": + _valid_resolutions = [1, 2] + elif space == "MNInonlinear2009cAsym": + _valid_resolutions = [2] + else: + raise_error( + f"The parameter `space` ({space}) for 3T needs to be one of " + f"the following: {_valid_spaces}" + ) + elif magneticfield == "7T": + _valid_resolutions = [1.6] + if space != "MNI6thgeneration": + raise_error( + f"The parameter `space` ({space}) for 7T needs to be " + f"MNI6thgeneration" + ) + else: + raise_error( + f"The parameter `magneticfield` ({magneticfield}) needs to be " + f"one of the following: 3T or 7T" + ) + + resolution = _closest_resolution(resolution, _valid_resolutions) + + # define file names + if magneticfield == "3T": + atlas_fname_base_3T = ( + atlas_dir / "Tian2020MSA_v1.1" / "3T" / "Subcortex-Only" + ) + atlas_lname = atlas_fname_base_3T / ( + f"Tian_Subcortex_S{scale}_3T_label.txt" + ) + if space == "MNI6thgeneration": + atlas_fname = atlas_fname_base_3T / ( + f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz" + ) + if resolution == 1: + atlas_fname = ( + atlas_fname_base_3T + / f"Tian_Subcortex_S{scale}_{magneticfield}_1mm.nii.gz" + ) + elif space == "MNInonlinear2009cAsym": + space = "2009cAsym" + atlas_fname = atlas_fname_base_3T / ( + f"Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz" + ) + elif magneticfield == "7T": + atlas_fname_base_7T = atlas_dir / "Tian2020MSA_v1.1" / "7T" + atlas_fname_base_7T.mkdir(exist_ok=True, parents=True) + atlas_fname = ( + atlas_dir + / "Tian2020MSA_v1.1" + / f"{magneticfield}" + / (f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz") + ) + # define 7T labels (b/c currently no labels file available for 7T) + scale7Trois = {1: 16, 2: 34, 3: 54, 4: 62} + labels = [ + ("parcel_" + str(x)) for x in np.arange(1, scale7Trois[scale] + 1) + ] + atlas_lname = atlas_fname_base_7T / ( + f"Tian_Subcortex_S{scale}_7T_labelnumbering.txt" + ) + with open(atlas_lname, "w") as filehandle: + for listitem in labels: + filehandle.write("%s\n" % listitem) + logger.info( + "Currently there are no labels provided for the 7T Tian atlas. " + "A simple numbering scheme for distinction was therefore used." + ) + else: # pragma: no cover + raise_error("This should not happen. Please report this error.") + + # check existence of atlas + if not (atlas_fname.exists() and atlas_lname.exists()): + logger.info("At least one of the atlas files is missing, fetching.") + + url_basis = ( + "https://www.nitrc.org/frs/download.php/12012/Tian2020MSA_v1.1.zip" + ) + + logger.info(f"Downloading TIAN from {url_basis}") + with tempfile.TemporaryDirectory() as tmpdir: + atlas_download = requests.get(url_basis) + atlas_zip_fname = Path(tmpdir) / "Tian2020MSA_v1.1.zip" + with open(atlas_zip_fname, "wb") as f: + f.write(atlas_download.content) + with zipfile.ZipFile(atlas_zip_fname, "r") as zip_ref: + zip_ref.extractall(atlas_dir.as_posix()) + # clean after unzipping + if (atlas_dir / "__MACOSX").exists(): + shutil.rmtree((atlas_dir / "__MACOSX").as_posix()) + + labels = pd.read_csv(atlas_lname, sep=" ", header=None)[0].to_list() + + if not (atlas_fname.exists() and atlas_lname.exists()): + raise_error("There was a problem fetching the atlases.") + + labels = pd.read_csv(atlas_lname, sep=" ", header=None)[0].to_list() + + return atlas_fname, labels + + +def _retrieve_suit( + atlas_dir: Path, resolution: Optional[float], space: str = "MNI" +) -> Tuple[Path, List[str]]: + """Retrieve SUIT atlas. + + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : float, optional + The desired resolution of the atlas to load. If it is not available, + the closest resolution will be loaded. Preferably, use a resolution + higher than the desired one. By default, will load the highest one + (default None). Available resolutions for this atlas are 1mm and 2mm. + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). + + Returns + ------ + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `space` or if there is a problem + fetching the atlas. + + """ + logger.info("Atlas parameters:") + logger.info(f"\tspace: {space}") + + _valid_spaces = ["MNI", "SUIT"] # check validity of atlas parameters if space not in _valid_spaces: raise_error( - f'The parameter `space` ({space}) needs to be one of the ' - f'following: {_valid_spaces}') + f"The parameter `space` ({space}) needs to be one of the " + f"following: {_valid_spaces}" + ) # TODO: Validate this with Vera _valid_resolutions = [1] @@ -324,45 +675,52 @@ def _retrieve_suit(out_dir, resolution, space='MNI'): resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - atlas_fname = out_dir / 'SUIT' / ( - f'SUIT_{space}Space_{resolution}mm.nii') - atlas_lname = out_dir / 'SUIT' / ( - f'SUIT_{space}Space_{resolution}mm.tsv') + atlas_fname = ( + atlas_dir / "SUIT" / (f"SUIT_{space}Space_{resolution}mm.nii") + ) + atlas_lname = ( + atlas_dir / "SUIT" / (f"SUIT_{space}Space_{resolution}mm.tsv") + ) - # check existance of atlas + # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): - logger.info( - 'At least one of the atlas files is missing. ' - 'Fetching.') + atlas_fname.parent.mkdir(exist_ok=True, parents=True) + logger.info("At least one of the atlas files is missing, fetching.") url_basis = ( - 'https://github.com/DiedrichsenLab/cerebellar_atlases/blob' - '/master/Diedrichsen_2009/') - url_MNI = url_basis + 'atl-Anatom_space-MNI_dseg.nii' - url_SUIT = url_basis + 'atl-Anatom_space-SUIT_dseg.nii' - url_labels = url_basis + 'atl-Anatom.tsv' + "https://github.com/DiedrichsenLab/cerebellar_atlases/raw" + "/master/Diedrichsen_2009/" + ) + url_MNI = url_basis + "atl-Anatom_space-MNI_dseg.nii" + url_SUIT = url_basis + "atl-Anatom_space-SUIT_dseg.nii" + url_labels = url_basis + "atl-Anatom.tsv" - if space == 'MNI': - logger.info(f'Downloading {url_MNI}') + if space == "MNI": + logger.info(f"Downloading {url_MNI}") atlas_download = requests.get(url_MNI) - with open(atlas_fname, 'wb') as f: + with open(atlas_fname, "wb") as f: f.write(atlas_download.content) - elif space == 'SUIT': - logger.info(f'Downloading {url_SUIT}') + else: # if not MNI, then SUIT + logger.info(f"Downloading {url_SUIT}") atlas_download = requests.get(url_SUIT) - with open(atlas_fname, 'wb') as f: + with open(atlas_fname, "wb") as f: f.write(atlas_download.content) labels_download = requests.get(url_labels) labels = pd.read_csv( io.StringIO(labels_download.content.decode("utf-8")), - sep='\t', usecols=['name']) + sep="\t", + usecols=["name"], + ) - labels.to_csv(atlas_lname, sep='\t', index=False) - if not (atlas_fname.exists() and atlas_lname.exists()): - raise_error('There was a problem fetching the atlases.') + labels.to_csv(atlas_lname, sep="\t", index=False) + if ( + not atlas_fname.exists() and atlas_lname.exists() + ): # pragma: no cover + raise_error("There was a problem fetching the atlases.") - labels = pd.read_csv( - atlas_lname, sep='\t', usecols=['name'])['name'].to_list() + labels = pd.read_csv(atlas_lname, sep="\t", usecols=["name"])[ + "name" + ].to_list() return atlas_fname, labels diff --git a/junifer/data/tests/test_atlases.py b/junifer/data/tests/test_atlases.py new file mode 100644 index 000000000..fb6190f88 --- /dev/null +++ b/junifer/data/tests/test_atlases.py @@ -0,0 +1,483 @@ +"""Provide tests for atlas.""" + +# Authors: Federico Raimondo +# Vera Komeyer +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import List + +import pytest +from numpy.testing import assert_array_almost_equal, assert_array_equal + +from junifer.data.atlases import ( + _retrieve_atlas, + _retrieve_schaefer, + _retrieve_suit, + _retrieve_tian, + list_atlases, + load_atlas, + register_atlas, +) + + +def test_register_atlas_built_in_check() -> None: + """Test atlas registration check for built-in atlas.""" + with pytest.raises(ValueError, match=r"built-in atlas"): + register_atlas( + name="SUITxSUIT", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + overwrite=True, + ) + + +def test_list_atlases_incorrect() -> None: + """Test incorrect information check for list atlases.""" + atlases = list_atlases() + assert "testatlas" not in atlases + + +def test_register_atlas_already_registered() -> None: + """Test atlas registration check for already registered atlas.""" + # Register custom atlas + register_atlas( + name="testatlas", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + ) + # Try registering again + with pytest.raises(ValueError, match=r"already registered."): + register_atlas( + name="testatlas", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + ) + + +@pytest.mark.parametrize( + "name, atlas_path, atlas_labels, overwrite", + [ + ("testatlas_1", "testatlas_1.nii.gz", ["1", "2", "3"], True), + ("testatlas_2", "testatlas_2.nii.gz", ["1", "2", "6"], True), + ("testatlas_3", Path("testatlas_3.nii.gz"), ["1", "2", "6"], True), + ], +) +def test_register_atlas( + name: str, + atlas_path: str, + atlas_labels: List[str], + overwrite: bool, +) -> None: + """Test atlas registration. + + Parameters + ---------- + name : str + The parametrized atlas name. + atlas_path : str or pathlib.Path + The parametrized atlas path. + atlas_labels : list of str + The parametrized atlas labels. + overwrite : bool + The parametrized atlas overwrite value. + + """ + # Register custom atlas + register_atlas( + name=name, + atlas_path=atlas_path, + atl_labels=atlas_labels, + overwrite=overwrite, + ) + # List available atlas and check registration + atlases = list_atlases() + assert name in atlases + # Load registered atlas + _, lbl, fname = load_atlas(name=name, path_only=True) + # Check values for registered atlas + assert lbl == atlas_labels + assert fname.name == f"{name}.nii.gz" + + +@pytest.mark.parametrize( + "atlas_name", + [ + "SUITxSUIT", + "SUITxMNI", + "Schaefer100x7", + "Schaefer100x17", + "TianxS1x7TxMNI6thgeneration", + "TianxS3x3TxMNI6thgeneration", + "TianxS4x3TxMNInonlinear2009cAsym", + ], +) +def test_list_atlases_correct(atlas_name: str) -> None: + """Test correct information check for list atlases. + + Parameters + ---------- + atlas_name : str + The parametrized atlas name. + + """ + atlases = list_atlases() + assert atlas_name in atlases + + +def test_load_atlas_incorrect() -> None: + """Test loading of invalid atlas.""" + with pytest.raises(ValueError, match=r"not found"): + load_atlas("wrongatlas") + + +def test_retrieve_atlas_incorrect() -> None: + """Test retrieval of invalid atlas.""" + with pytest.raises(ValueError, match=r"provided atlas name"): + _retrieve_atlas("wrongatlas") + + +# TODO: paramdtrize test +def test_schaefer_atlas(tmp_path: Path) -> None: + """Test Schaefer atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + atlases = list_atlases() + for n_rois in range(100, 1001, 100): + for t_net in [7, 17]: + t_name = f"Schaefer{n_rois}x{t_net}" + assert t_name in atlases + + # Define atlas file names + fname1 = "Schaefer2018_100Parcels_7Networks_order_FSLMNI152_1mm.nii.gz" + fname2 = "Schaefer2018_100Parcels_7Networks_order_FSLMNI152_2mm.nii.gz" + + # Load atlas + img, lbl, fname = load_atlas( + name="Schaefer100x7", atlas_dir=str(tmp_path.absolute()) + ) + # Check atlas values + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 100 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + + # Test with Path + img, lbl, fname = load_atlas(name="Schaefer100x7", atlas_dir=tmp_path) + # Load atlas + img2, lbl, fname = load_atlas( + name="Schaefer100x7", + atlas_dir=tmp_path, + resolution=3, + ) + # Check atlas values + assert fname.name == fname2 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=2.1, + ) + # Check atlas values + assert fname.name == fname2 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=1.99, + ) + # Check atlas values + assert fname.name == fname1 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=0.5, + ) + # Check atlas values + assert fname.name == fname1 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + + +def test_load_atlas_schaefer() -> None: + """Test Schaefer atlas loading.""" + img, lbl, fname = load_atlas(name="Schaefer100x7") + assert img is not None + home_dir = Path().home() / "junifer" / "data" / "atlas" + assert home_dir in fname.parents + + +def test_retrieve_schaefer_incorrect_n_rois(tmp_path: Path) -> None: + """Test retrieve schaefer with incorrect n_rois. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `n_rois`"): + _retrieve_schaefer( + atlas_dir=tmp_path, resolution=1, n_rois=101, yeo_networks=7 + ) + + +def test_retrieve_schaefer_incorrect_yeo_networks(tmp_path: Path) -> None: + """Test retrieve schaefer with incorrect yeo_networks. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"): + _retrieve_schaefer( + atlas_dir=tmp_path, resolution=1, n_rois=100, yeo_networks=8 + ) + + +# TODO: parametrize test +def test_suit(tmp_path: Path) -> None: + """Test SUIT atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + atlases = list_atlases() + assert "SUITxSUIT" in atlases + assert "SUITxMNI" in atlases + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxSUIT", atlas_dir=tmp_path) + fname1 = "SUIT_SUITSpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxSUIT", atlas_dir=tmp_path) + fname1 = "SUIT_SUITSpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxMNI", atlas_dir=tmp_path) + fname1 = "SUIT_MNISpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + + +def test_retrieve_suit_incorrect_space(tmp_path: Path) -> None: + """Test retrieve suit with incorrect space. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `space`"): + _retrieve_suit(atlas_dir=tmp_path, resolution=1, space="wrong") + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 32), + (3, 50), + (4, 54), + ], +) +def test_tian_3T_6thgeneration( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x3TxMNI6thgeneration" in atlases + assert "TianxS2x3TxMNI6thgeneration" in atlases + assert "TianxS3x3TxMNI6thgeneration" in atlases + assert "TianxS4x3TxMNI6thgeneration" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNI6thgeneration", + atlas_dir=tmp_path, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T_1mm.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNI6thgeneration", + atlas_dir=tmp_path, + resolution=2, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) # type: ignore + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 32), + (3, 50), + (4, 54), + ], +) +def test_tian_3T_nonlinear2009cAsym( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS2x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS3x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS4x3TxMNInonlinear2009cAsym" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNInonlinear2009cAsym", + atlas_dir=tmp_path, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T_2009cAsym.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) # type: ignore + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 34), + (3, 54), + (4, 62), + ], +) +def test_tian_7T_6thgeneration( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x7TxMNI6thgeneration" in atlases + assert "TianxS2x7TxMNI6thgeneration" in atlases + assert "TianxS3x7TxMNI6thgeneration" in atlases + assert "TianxS4x7TxMNI6thgeneration" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x7TxMNI6thgeneration", atlas_dir=tmp_path + ) + fname1 = f"Tian_Subcortex_S{scale}_7T.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_almost_equal( + img.header["pixdim"][1:4], [1.6, 1.6, 1.6] + ) # type: ignore + + +def test_retrieve_tian_incorrect_space(tmp_path: Path) -> None: + """Test retrieve tian with incorrect space. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `space`"): + _retrieve_tian( + atlas_dir=tmp_path, + resolution=1, + scale=1, + space="wrong", + ) + + +def test_retrieve_tian_incorrect_magneticfield(tmp_path: Path) -> None: + """Test retrieve tian with incorrect magneticfield. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): + _retrieve_tian( + atlas_dir=tmp_path, + resolution=1, + scale=1, + magneticfield="wrong", + ) diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index e9866a1be..5975d9b5d 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -1,4 +1,12 @@ +"""Provide imports for datagrabber sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from .base import BIDSDataladDataGrabber, DataladDataGrabber, BIDSDataGrabber \ No newline at end of file + +from .base import BaseDataGrabber +from .datalad_base import DataladDataGrabber +from .hcp import DataladHCP1200, HCP1200 +from .multiple import MultipleDataGrabber +from .pattern import PatternDataGrabber +from .pattern_datalad import PatternDataladDataGrabber diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 8c828c875..8f7065e46 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -1,294 +1,158 @@ +"""Provide abstract base class for datagrabber.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from pathlib import Path -import tempfile -import datalad.api as dl from abc import ABC, abstractmethod +from pathlib import Path +from typing import Dict, Iterator, List, Tuple, Union -from ..api.decorators import register_datagrabber -from ..utils.logging import logger, raise_error - - -def _validate_types(types): - """ - Validate the types - """ - if not isinstance(types, list): - raise_error("types must be a list", TypeError) - if any(not isinstance(x, str) for x in types): - raise_error("types must be a list of strings", TypeError) - - -def _validate_patterns(types, patterns): - """ - Validate the patterns. - """ - _validate_types(types) - if not isinstance(patterns, dict): - raise_error("patterns must be a dict", TypeError) - if len(types) != len(patterns): - raise_error("types and patterns must have the same length", ValueError) - - if any(x not in patterns for x in types): - raise_error("patterns must contain all types", ValueError) +from ..utils import logger, raise_error +from .utils import validate_types class BaseDataGrabber(ABC): - """Base DataGrabber class (abstract). + """Abstract base class for datagrabber. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + Parameters + ---------- + types : list of str + The types of data to be grabbed. + datadir : str or pathlib.Path + The directory where the data is / will be stored. Attributes ---------- - datadir - types : list - List of data types to be grabbed. + datadir : pathlib.Path + The directory where the data is / will be stored. - Methods - ------- - get_elements : List - Returns a list of elements that can be grabbed. The elements can be - strings, tuples or any object that will be then used as a key to - index the datagrabber - __getitem__(element) : dict[str -> Path] - Returns a dictionary of paths for each type of data required for the - specified element. Use the element as a key to index the datagrabber. - __enter__() : self - Returns the object itself. Can be overridden by subclasses. - __exit__() : None - Does nothing. Can be overridden by subclasses to clean up after - `__enter__` """ - def __init__(self, types, datadir): - """Initialize a BaseDataGrabber object. - Parameters - ---------- - types : list of str - The types of data to be grabbed. - datadir : str or Path - That directory where the data is/will be stored. - """ - _validate_types(types) + def __init__(self, types: List[str], datadir: Union[str, Path]) -> None: + """Initialize the class.""" + # Validate types + validate_types(types) + # Convert str to Path if not isinstance(datadir, Path): datadir = Path(datadir) + logger.debug("Initializing BaseDataGrabber") + logger.debug(f"\t_datadir = {datadir}") + logger.debug(f"\ttypes = {types}") self._datadir = datadir self.types = types - @property - def datadir(self): - """ - Returns - ------- - Path to the data directory. Implemented as a property, can be - overridden by subclasses. - """ - return self._datadir - - def __iter__(self): - """Iterate over elements in the datagrabber. + def __iter__(self) -> Iterator: + """Enable iterable support. Yields ------ - element : object + object An element that can be indexed by the datagrabber. + """ for elem in self.get_elements(): yield elem - @abstractmethod - def __getitem__(self, element): - raise NotImplementedError('__getitem__ not implemented') - - @abstractmethod - def get_elements(self): - raise NotImplementedError('get_elements not implemented') - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc_value, exc_traceback): - pass - - -@register_datagrabber -class BIDSDataGrabber(BaseDataGrabber): - """BIDS DataGrabber class. Implements a DataGrabber that understands BIDS - database format. - - Attributes - ---------- - datadir - types : list - List of data types to be grabbed. - patterns : dict[str -> str] - Patterns for each type of data. - - Methods - ------- - get_elements: list[str] - Returns a list of elements that can be grabbed. Each element is a - subject in the BIDS database. - __getitem__(str): dict[str -> Path] - Returns a dictionary of paths for each type of data required for the - specified element. Each occurrence of the string `{subject}` is - replaced by the indexed element - """ - def __init__(self, types=None, patterns=None, **kwargs): - """Initialize a BaseDataGrabber object. + # TODO: element does nothing, check + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]: + """Enable indexing support. Parameters ---------- - types : list of str - The types of data to be grabbed. - patterns : dict[str -> str] - Patterns for each type of data. The keys are the types and the - values are the patterns. Each occurrence of the string `{subject}` - in the pattern will be replaced by the indexed element. - datadir : str or Path - That directory where the data is/will be stored. - """ - _validate_patterns(types, patterns) - super().__init__(types=types, **kwargs) - self.patterns = patterns - - def get_elements(self): - """Get all the elements in the BIDS database + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". Returns ------- - elems : list[str] - List of all the elements in the database root directory - """ - elems = [x.name for x in self.datadir.iterdir() if x.is_dir()] - return elems - - def __getitem__(self, element): - """Index one element in the BIDS database. - - Each occurrence of the string `{subject}` is replaced by the indexed - element. - - Parameters - ---------- - element : str - The element to be indexed. - Returns - ------- - out : dict[str -> Path] + dict Dictionary of paths for each type of data required for the specified element. - """ - out = {} - for t_type in self.types: - t_pattern = self.patterns[t_type] # type: ignore - t_replace = t_pattern.replace('{subject}', element) - t_out = self.datadir / element / t_replace - out[t_type] = dict(path=t_out) + """ + logger.info(f"Getting element {element}") + out = {} + out["meta"] = {"datagrabber": self.get_meta()} return out - -@register_datagrabber -class DataladDataGrabber(BaseDataGrabber): - """ - Datalad DataGrabber class. Implements a DataGrabber that gets data from - a datalad sibling. - - Attributes - ---------- - datadir - uri : str - URI of the datalad sibling. - - Methods - ------- - install: - Installs (clones) the datalad dataset into the datadir. This method - is called automatically when the datagrabber is used within a `with` - statement. - remove: - Remove the datalad dataset from the datadir. This method is called - automatically when the datagrabber is used within a `with` statement. - - Note - ---- - By itself, this class is still abstract as the `__getitem__` method relies - on the parent class `BaseDataGrabber.__getitem__` which is not yet - implemented. This class is intended to be used as a superclass of a class - with multiple inheritance. See :class:`BIDSDataladDataGrabber` for a - concrete class implementation. - - """ - def __init__(self, rootdir='.', datadir=None, uri=None, **kwargs): - """Initialize a DataladDataGrabber object. - - Parameters - ---------- - rootdir : str or Path - The path within the datalad dataset to the root directory. - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - uri : str - URI of the datalad sibling. - """ - if uri is None: - raise_error('uri must be provided', ValueError) - if datadir is None: - logger.warning('datadir is None, creating a temporary directory') - datadir = tempfile.mkdtemp() - logger.info(f'datadir set to {datadir}') - super().__init__(datadir=datadir, **kwargs) - self.uri = uri - self._rootdir = rootdir - - def __enter__(self): - self.install() + def __enter__(self) -> "BaseDataGrabber": + """Context entry.""" return self + def __exit__(self, exc_type, exc_value, exc_traceback) -> None: + """Context exit.""" + return None + + def get_types(self) -> List[str]: + """Get types. + + Returns + ------- + list of str + The types of data to be grabbed. + + """ + return self.types.copy() + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith("_"): + t_meta[k] = v + return t_meta + + # TODO: what is the final functionality? + def get_element_keys(self) -> str: + """Get element keys. + + Returns + ------- + str + + """ + return "element" + @property - def datadir(self): - return super().datadir / self._rootdir + def datadir(self) -> Path: + """Get data directory path. - def install(self): - """Install the datalad dataset into the datadir.""" - self.dataset = dl.install( # type: ignore - self._datadir, source=self.uri) + Returns + ------- + pathlib.Path + Path to the data directory. Can be overridden by subclasses. - def __exit__(self, exc_type, exc_value, exc_traceback): - self.remove() + """ + return self._datadir - def remove(self): - """Remove the datalad dataset from the datadir.""" - self.dataset.remove(recursive=True) + @abstractmethod + def get_elements(self) -> List: + """Get elements. - def _dataset_get(self, out): - for _, v in out.items(): - self.dataset.get(v['path']) + Returns + ------- + list + List of elements that can be grabbed. The elements can be strings, + tuples or any object that will be then used as a key to index the + datagrabber. - def __getitem__(self, element): - """Index one element in the Datalad database. It will first obtain - the paths from the parent class and then `datalad get` each of the - files.""" - out = super().__getitem__(element) - self._dataset_get(out) - return out - - -class BIDSDataladDataGrabber(DataladDataGrabber, BIDSDataGrabber): - """BIDS Datalad DataGrabber class. - Implements a DataGrabber that gets data from a datalad sibling which - follows a BIDS format. - - See Also - -------- - DataladDataGrabber - BIDSDataGrabber - - """ - def __init__(self, types=None, patterns=None, **kwargs): - _validate_patterns(types, patterns) - super().__init__(types=types, patterns=patterns, **kwargs) - self.patterns = patterns + """ + raise_error( + msg="Concrete classes need to implement get_elements().", + klass=NotImplementedError, + ) diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py new file mode 100644 index 000000000..a0b004323 --- /dev/null +++ b/junifer/datagrabber/datalad_base.py @@ -0,0 +1,167 @@ +"""Provide abstract base class for datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import tempfile +from pathlib import Path +from typing import Dict, Optional, Tuple, Union + +import datalad.api as dl + +from ..api.decorators import register_datagrabber +from ..utils import logger +from .base import BaseDataGrabber +from .utils import raise_error + + +@register_datagrabber +class DataladDataGrabber(BaseDataGrabber): + """Abstract base class for data fetching via Datalad. + + Defines a DataGrabber that gets data from a datalad sibling. + + Parameters + ---------- + rootdir : str or Path, optional + The path within the datalad dataset to the root directory + (default "."). + datadir : str or Path, optional + That directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + uri : str, optional + URI of the datalad sibling (default None). + **kwargs + Keyword arguments passed to superclass. + + Methods + ------- + install: + Installs (clones) the datalad dataset into the `datadir`. This method + is called automatically when the datagrabber is used within a context. + remove: + Removes the datalad dataset from the `datadir`. This method is called + automatically when the datagrabber is used within a context. + + Notes + ----- + By itself, this class is still abstract as the `__getitem__` method relies + on the parent class `BaseDataGrabber.__getitem__` which is not yet + implemented. This class is intended to be used as a superclass of a class + with multiple inheritance. + + See Also + -------- + BaseDataGrabber + + """ + + def __init__( + self, + rootdir: Union[str, Path] = ".", + datadir: Union[str, Path, None] = None, + uri: Optional[str] = None, + **kwargs, + ): + """Initialize the class.""" + if datadir is None: + logger.warning("`datadir` is None, creating a temporary directory") + # Create temporary directory + datadir = tempfile.mkdtemp() + logger.info(f"`datadir` set to {datadir}") + # TODO: uri can be converted to a positional argument + if uri is None: + raise_error("`uri` must be provided") + + super().__init__(datadir=datadir, **kwargs) + logger.debug("Initializing DataladDataGrabber") + logger.debug(f"\turi = {uri}") + logger.debug(f"\t_rootdir = {rootdir}") + self.uri = uri + self._rootdir = rootdir + + @property + def datadir(self) -> Path: + """Get data directory path.""" + return super().datadir / self._rootdir + + def _dataset_get(self, out: Dict) -> Dict: + """Get the dataset found from the path in `out`. + + Parameters + ---------- + out : dict + The dictionary from which path need to be searched. + + Returns + ------- + dict + The modified dictionary with version appended. + + """ + for _, v in out.items(): + if "path" in v: + logger.debug(f"Getting {v['path']}") + self._dataset.get(v["path"]) + logger.debug("Get done") + + # append the version of the dataset + out["meta"]["datagrabber"][ + "dataset_commit_id" + ] = self._dataset.repo.get_hexsha( + self._dataset.repo.get_corresponding_branch() + ) + return out + + def install(self) -> None: + """Install the datalad dataset into the datadir.""" + logger.debug(f"Installing dataset {self.uri} to {self._datadir}") + self._dataset = dl.install( # type: ignore because of datalad + self._datadir, source=self.uri + ) + logger.debug("Dataset installed") + + def remove(self): + """Remove the datalad dataset from the datadir.""" + self._dataset.remove(recursive=True) + + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Implement single element indexing in the Datalad database. + + It will first obtain the paths from the parent class and then + `datalad get` each of the files. + + This method only works with multiple inheritance. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + out = super().__getitem__(element) + out = self._dataset_get(out) + return out + + def __enter__(self): + """Implement context entry.""" + self.install() + return self + + def __exit__(self, exc_type, exc_value, exc_traceback): + """Implement context exit.""" + logger.debug("Removing dataset") + self.remove() + logger.debug("Dataset removed") diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py new file mode 100644 index 000000000..9db59031b --- /dev/null +++ b/junifer/datagrabber/hcp.py @@ -0,0 +1,201 @@ +"""Provide concrete implementations for HCP data access.""" + +from itertools import product +from pathlib import Path +from typing import Dict, List, Tuple, Union + +from junifer.datagrabber.datalad_base import DataladDataGrabber + +from ..api.decorators import register_datagrabber +from ..utils import raise_error +from .pattern import PatternDataGrabber + + +@register_datagrabber +class HCP1200(PatternDataGrabber): + """Concrete implementation for pattern-based data fetching of HCP1200. + + Parameters + ---------- + datadir : str or Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", + "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional + HCP task sessions. If None, all available task sessions are selected + (default None). + phase_encodings : {"LR", "RL"} or list of the options, optional + HCP phase encoding directions. If None, both will be used + (default None). + **kwargs + Keyword arguments passed to superclass. + + """ + + def __init__( + self, + datadir: Union[str, Path, None] = None, + tasks: Union[str, List[str], None] = None, + phase_encodings: Union[str, List[str], None] = None, + **kwargs, + ) -> None: + """Initialize the class.""" + # All tasks + all_tasks = [ + "REST1", + "REST2", + "SOCIAL", + "WM", + "RELATIONAL", + "EMOTION", + "LANGUAGE", + "GAMBLING", + "MOTOR", + ] + # Set default tasks + if tasks is None: + self.tasks: List[str] = all_tasks + # Convert single task into list + else: + if not isinstance(tasks, List): + tasks = [tasks] + # Check for invalid task(s) + for task in tasks: + if task not in all_tasks: + raise_error( + f"'{task}' is not a valid HCP-YA fMRI task input. " + f"Valid task values can be any or all of {all_tasks}." + ) + self.tasks: List[str] = tasks + # All phase encodings + all_phase_encodings = ["LR", "RL"] + # Set phase encodings + if phase_encodings is None: + phase_encodings = all_phase_encodings + # Convert single phase encoding into list + if isinstance(phase_encodings, str): + phase_encodings = [phase_encodings] + # Check for invalid phase encoding(s) + for pe in phase_encodings: + if pe not in all_phase_encodings: + raise_error( + f"'{pe}' is not a valid HCP-YA phase encoding. " + "Valid phase encoding can be any or all of " + f"{all_phase_encodings}." + ) + + # The types of data + types = ["BOLD"] + # The patterns + patterns = { + "BOLD": ( + "{subject}/MNINonLinear/Results/" + "{task}_{phase_encoding}/" + "{task}_{phase_encoding}_hp2000_clean.nii.gz" + ) + } + # The replacements + replacements = ["subject", "task", "phase_encoding"] + super().__init__( + types=types, + datadir=datadir, + patterns=patterns, + replacements=replacements, + ) + self.phase_encodings = phase_encodings + + def __getitem__(self, element: Tuple[str, str, str]) -> Dict[str, Path]: + """Index one element in the dataset. + + Parameters + ---------- + element : triple of str + The element to be indexed. First element in the tuple is the + subject, second element is the task, third element is the + phase encoding direction. + + Returns + ------- + out : dict + Dictionary of paths for each type of data required for the + specified element. + + """ + sub, task, phase_encoding = element + + # Resting task + if "REST" in task: + new_task = f"rfMRI_{task}" + else: + new_task = f"tfMRI_{task}" + + out = super().__getitem__((sub, new_task, phase_encoding)) + out["meta"]["element"] = { + "subject": sub, + "task": task, + "phase_encoding": phase_encoding, + } + return out + + def get_elements(self) -> List: + """Implement fetching list of subjects in the dataset. + + Returns + ------- + elements : list of str + The list of subjects in the dataset. + + """ + subjects = [x.name for x in self.datadir.iterdir() if x.is_dir()] + elems = [] + for subject, task, phase_encoding in product( + subjects, self.tasks, self.phase_encodings + ): + elems.append((subject, task, phase_encoding)) + + return elems + + +@register_datagrabber +class DataladHCP1200(DataladDataGrabber, HCP1200): + """Concrete implementation for datalad-based data fetching of HCP1200. + + Parameters + ---------- + datadir : str or Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", + "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional + HCP task sessions. If None, all available task sessions are selected + (default None). + phase_encodings : {"LR", "RL"} or list of the options, optional + HCP phase encoding directions. If None, both will be used + (default None). + **kwargs + Keyword arguments passed to superclass. + + """ + + def __init__( + self, + datadir: Union[str, Path, None] = None, + tasks: Union[str, List[str], None] = None, + phase_encodings: Union[str, List[str], None] = None, + **kwargs, + ) -> None: + """Initialize the class.""" + uri = ( + "https://github.com/datalad-datasets/" + "human-connectome-project-openaccess.git" + ) + rootdir = "HCP1200" + super().__init__( + datadir=datadir, + tasks=tasks, + phase_encodings=phase_encodings, + uri=uri, + rootdir=rootdir, + ) diff --git a/junifer/datagrabber/multiple.py b/junifer/datagrabber/multiple.py new file mode 100644 index 000000000..606fc82bc --- /dev/null +++ b/junifer/datagrabber/multiple.py @@ -0,0 +1,111 @@ +"""Provide abstract base class for multiple source datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import Dict, List, Tuple, Union + +from .base import BaseDataGrabber + + +class MultipleDataGrabber(BaseDataGrabber): + """Datagrabber class for data fetching from multiple sources. + + Defines a DataGrabber which can be used to fetch data from multiple + datagrabbers. + + Parameters + ---------- + datagrabbers : list of datagrabbers + The datagrabbers to use to fetch data using. + **kwargs + Keyword arguments passed to superclass. + + """ + + def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None: + """Initialize the class.""" + # TODO: Check datagrabbers consistency + # - same element keys + # - no overlapping types + self._datagrabbers = datagrabbers + + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Implement indexing. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + out = {} + for dg in self._datagrabbers: + t_out = dg[element] + out.update(t_out) + return out + + def __enter__(self) -> "BaseDataGrabber": + """Implement context entry.""" + for dg in self._datagrabbers: + dg.__enter__() + return self + + def __exit__(self, exc_type, exc_value, exc_traceback) -> None: + """Implement context exit.""" + for dg in self._datagrabbers: + dg.__exit__(exc_type, exc_value, exc_traceback) + + def get_elements(self) -> List: + """Get elements. + + Returns + ------- + elements : list + The list of elements that can be grabbed in the dataset. It + corresponds to the elements that are present in all the + related datagrabbers. + """ + all_elements = [dg.get_elements() for dg in self._datagrabbers] + elements = set(all_elements[0]) + for s in all_elements[1:]: + elements.intersection_update(s) + return list(elements) + + def get_types(self) -> List[str]: + """Get types. + + Returns + ------- + list of list of str + The types of data to be grabbed. + + """ + types = [x for dg in self._datagrabbers for x in dg.get_types()] + return types + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + t_meta["datagrabbers"] = [dg.get_meta() for dg in self._datagrabbers] + return t_meta diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py new file mode 100644 index 000000000..e32dc45af --- /dev/null +++ b/junifer/datagrabber/pattern.py @@ -0,0 +1,204 @@ +"""Provide concrete implementation for pattern-based datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import re +from typing import Dict, List, Tuple, Union + +from ..api.decorators import register_datagrabber +from ..utils import logger, raise_error +from .base import BaseDataGrabber +from .utils import validate_patterns, validate_replacements + + +@register_datagrabber +class PatternDataGrabber(BaseDataGrabber): + """Concrete implementation for data fetching using patterns. + + Implements a DataGrabber that understands patterns to grab data. + + Parameters + ---------- + types : list of str + The types of data to be grabbed (default None). + patterns : dict + Patterns for each type of data as a dictionary. The keys are the types + and the values are the patterns. Each occurrence of the string + `{subject}` in the pattern will be replaced by the indexed element. + replacements: list of str + Replacements in the patterns for each item in the "element" tuple. + datadir : str or pathlib.Path + The directory where the data is / will be stored. + **kwargs + Keyword arguments passed to superclass. + + See Also + -------- + BaseDataGrabber + + """ + + def __init__( + self, + types: List[str], + patterns: Dict[str, str], + replacements: List[str], + **kwargs, + ) -> None: + """Initialize the class.""" + # Validate patterns + validate_patterns(types=types, patterns=patterns) + + if not isinstance(replacements, list): + replacements = [replacements] + # Validate replacements + validate_replacements(replacements=replacements, patterns=patterns) + + super().__init__(types=types, **kwargs) + logger.debug("Initializing PatternDataGrabber") + logger.debug(f"\tpatterns = {patterns}") + logger.debug(f"\treplacements = {replacements}") + self.patterns = patterns + self.replacements = replacements + + def _replace_patterns_regex(self, pattern: str) -> Tuple[str, str]: + """Replace the patterns in `pattern` with the named groups. + + It allows elements to be obtained from the filesystem. + + Parameters + ---------- + pattern : str + The pattern to be replaced. + + Returns + ------- + re_pattern : str + The regular expression with the named groups. + glob_pattern : str + The search pattern to be used with glob. + + """ + re_pattern = pattern + glob_pattern = pattern + for t_r in self.replacements: + # Replace the first of each with a named group definition + re_pattern = re_pattern.replace(f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1) + + for t_r in self.replacements: + # Replace the second appearance of each with the named group + # back reference + re_pattern = re_pattern.replace(f"{{{t_r}}}", f"(?P={t_r})") + + for t_r in self.replacements: + glob_pattern = glob_pattern.replace(f"{{{t_r}}}", "*") + return re_pattern, glob_pattern + + def _replace_patterns_glob(self, element: Tuple, pattern: str) -> str: + """Replace patterns with the element so it can be globbed. + + Parameters + ---------- + element : tuple + The element to be used in the replacement. + pattern : str + The pattern to be replaced. + + Returns + ------- + str + The pattern with the element replaced. + + """ + if len(element) != len(self.replacements): + raise_error( + f"The element length must be {len(self.replacements)}, " + f"indicating {self.replacements}." + ) + to_replace = dict(zip(self.replacements, element)) + return pattern.format(**to_replace) + + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]: + """Implement single element indexing in the database. + + Each occurrence of the strings in "replacements" is replaced by the + corresponding item in the element tuple. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of dictionaries for each type of data required for the + specified element. + + """ + out = super().__getitem__(element) + if not isinstance(element, tuple): + element = (element,) + for t_type in self.types: + t_pattern = self.patterns[t_type] + t_replace = self._replace_patterns_glob(element, t_pattern) + if "*" in t_replace: + t_matches = list(self.datadir.glob(t_replace)) + if len(t_matches) > 1: + raise_error( + f"More than one file matches for {element} / {t_type}:" + f" {t_matches}" + ) + elif len(t_matches) == 0: + raise_error(f"No file matches for {element} / {t_type}") + t_out = t_matches[0] + else: + t_out = self.datadir / t_replace + out[t_type] = {"path": t_out} + # Meta here is element and types + out["meta"]["element"] = dict(zip(self.replacements, element)) + return out + + def get_elements(self) -> List: + """Implement fetching list of elements in the dataset. + + It will use regex to search for "replacements" in the "patterns" and + return the intersection of the results for each type i.e., build a + list of elements that have all the required types. + + Returns + ------- + elements : list + The list of elements that can be grabbed in the dataset. Each + element is a subject in the BIDS database. + + """ + elements = None + for t_type in self.types: + types_element = set() + # Get the pattern + t_pattern = self.patterns[t_type] + # Replace the pattern + re_pattern, glob_pattern = self._replace_patterns_regex(t_pattern) + for fname in self.datadir.glob(glob_pattern): + suffix = fname.relative_to(self.datadir).as_posix() + m = re.match(re_pattern, suffix) + if m is not None: + t_element = tuple(m.group(k) for k in self.replacements) + if len(self.replacements) == 1: + t_element = t_element[0] + types_element.add(t_element) + # TODO: does this make sense as elements is always None + if elements is None: + elements = types_element + else: + elements = elements.intersection(types_element) + if elements is None: + elements = set() + return list(elements) diff --git a/junifer/datagrabber/pattern_datalad.py b/junifer/datagrabber/pattern_datalad.py new file mode 100644 index 000000000..07e5e296d --- /dev/null +++ b/junifer/datagrabber/pattern_datalad.py @@ -0,0 +1,53 @@ +"""Provide base class for pattern-based datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..api.decorators import register_datagrabber +from .datalad_base import DataladDataGrabber +from .pattern import PatternDataGrabber +from .utils import validate_patterns + + +@register_datagrabber +class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): + """Base class for pattern-based data fetching via Datalad. + + Defines a DataGrabber that gets data from a datalad sibling, + interpreting patterns. + + Parameters + ---------- + types : list of str + The types of data to be grabbed. + patterns : dict, optional + Patterns for each type of data as a dictionary. The keys are the types + and the values are the patterns. Each occurrence of the string + `{subject}` in the pattern will be replaced by the indexed element + (default None). + **kwargs + Keyword arguments passed to superclass. + + See Also + -------- + DataladDataGrabber + PatternDataGrabber + + """ + + def __init__( + self, + types: List[str], + patterns: Dict[str, str], + **kwargs, + ) -> None: + """Initialize the class.""" + # Validate patterns + validate_patterns(types=types, patterns=patterns) + + super().__init__(types=types, patterns=patterns, **kwargs) + self.patterns = patterns diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base.py index 9d80054e1..91d622d62 100644 --- a/junifer/datagrabber/tests/test_base.py +++ b/junifer/datagrabber/tests/test_base.py @@ -1,76 +1,43 @@ +"""Provide tests for base.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -import pytest + from pathlib import Path -from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber + +import pytest + +from junifer.datagrabber.base import BaseDataGrabber -def test_BIDSDataGrabber(): - """Test BIDSDataGrabber""" - with pytest.raises(TypeError, match=r"types must be a list"): - BIDSDataGrabber(datadir='/tmp', types='wrong', - patterns=dict(wrong='pattern')) - - with pytest.raises(TypeError, match=r"must be a list of strings"): - BIDSDataGrabber(datadir='/tmp', types=[1, 2, 3], - patterns={'1': 'pattern', '2': 'pattern', - '3': 'pattern'}) - - datagrabber = BIDSDataGrabber( - datadir='/tmp/data', types=['func', 'anat'], - patterns=dict(func='pattern1', anat='pattern2')) - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] - - datagrabber = BIDSDataGrabber( - datadir=Path('/tmp/data'), types=['func', 'anat'], - patterns=dict(func='pattern1', anat='pattern2')) - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] - - with pytest.raises(TypeError, match=r"patterns must be a dict"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns='wrong') - - with pytest.raises(ValueError, - match=r"patterns must have the same length"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern'}) - - with pytest.raises(ValueError, match=r"patterns must contain all types"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern', 'func': 'pattern'}) +def test_BaseDataGrabber_abstractness() -> None: + """Test BaseDataGrabber is abstract base class.""" + with pytest.raises(TypeError, match=r"abstract"): + BaseDataGrabber(datadir="/tmp", types=["func"]) # type: ignore -def test_BIDSDataladDataGrabber(): - """Test BIDSDataladDataGrabber""" - types = ['T1w', 'bold'] - patterns = { - 'T1w': 'anat/{subject}_T1w.nii.gz', - 'bold': 'func/{subject}_task-rest_bold.nii.gz' - } +def test_BaseDataGrabber() -> None: + """Test BaseDataGrabber.""" + # Create concrete class. + class MyDataGrabber(BaseDataGrabber): + def __getitem__(self, element): + return super().__getitem__(element) - with pytest.raises(ValueError, match=r"uri must be provided"): - BIDSDataladDataGrabber(datadir=None, types=types, patterns=patterns) + def get_elements(self): + return super().get_elements() - repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids' - rootdir = 'example_bids' + dg = MyDataGrabber(datadir="/tmp", types=["func"]) + elem = dg["elem"] + assert "meta" in elem + assert "datagrabber" in elem["meta"] + assert "class" in elem["meta"]["datagrabber"] + assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"] - with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, - types=types, patterns=patterns) as dg: - subs = [x for x in dg] - expected_subs = [f'sub-{i:02d}' for i in range(1, 10)] - assert set(subs) == set(expected_subs) + with pytest.raises(NotImplementedError): + dg.get_elements() - for elem in dg: - t_sub = dg[elem] - assert 'path' in t_sub['T1w'] - assert t_sub['T1w']['path'] == \ - (dg.datadir / f'{elem}/anat/{elem}_T1w.nii.gz') - assert 'path' in t_sub['bold'] - assert t_sub['bold']['path'] == \ - (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') - - with open(t_sub['T1w']['path'], 'r') as f: - assert f.readlines()[0] == 'placeholder' + with dg: + assert dg.datadir == Path("/tmp") + assert dg.types == ["func"] diff --git a/junifer/datagrabber/tests/test_datalad_base.py b/junifer/datagrabber/tests/test_datalad_base.py new file mode 100644 index 000000000..2197fc097 --- /dev/null +++ b/junifer/datagrabber/tests/test_datalad_base.py @@ -0,0 +1,22 @@ +"""Provide tests for datalad_base.""" + +# Authors: Synchon Mandal +# License: AGPL + +import pytest + +from junifer.datagrabber.datalad_base import DataladDataGrabber + + +def test_datalad_base_abstractness() -> None: + """Test datalad base is abstract.""" + with pytest.raises(TypeError, match=r"abstract"): + DataladDataGrabber() + + +# def test_datalad_base_missing_uri() -> None: +# """Test proper check of missing URI in datalad base initialization.""" +# with pytest.raises(ValueError, match=r"`uri` must be provided"): +# DataladDataGrabber( + +# ) diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py new file mode 100644 index 000000000..5d591c23a --- /dev/null +++ b/junifer/datagrabber/tests/test_multiple.py @@ -0,0 +1,94 @@ +"""Provide tests for multiple.""" + +# Authors: Federico Raimondo +# License: AGPL + +from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber + + +_testing_dataset = { + "example_bids": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids", + "id": "e2ce149bd723088769a86c72e57eded009258c6b", + }, + "example_bids_ses": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "id": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + }, +} + + +def test_multiple() -> None: + """Test a multiple datagrabber.""" + repo_uri = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + pattern1 = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + pattern2 = { + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=["T1w"], + patterns=pattern1, + replacements=replacements, + ) + + dg2 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=["bold"], + patterns=pattern2, + replacements=replacements, + ) + + dg = MultipleDataGrabber([dg1, dg2]) + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 3) + for i in range(1, 10) + ] + + with dg: + subs = [x for x in dg] + assert set(subs) == set(expected_subs) + + +def test_multiple_no_intersection() -> None: + """Test a multiple datagrabber without intersection (0 elements).""" + repo_uri1 = _testing_dataset["example_bids"]["uri"] + repo_uri2 = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + pattern1 = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + pattern2 = { + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri1, + types=["T1w"], + patterns=pattern1, + replacements=replacements, + ) + + dg2 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri2, + types=["bold"], + patterns=pattern2, + replacements=replacements, + ) + + dg = MultipleDataGrabber([dg1, dg2]) + expected_subs = set() + with dg: + subs = [x for x in dg] + assert set(subs) == set(expected_subs) diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py new file mode 100644 index 000000000..530b7e447 --- /dev/null +++ b/junifer/datagrabber/tests/test_pattern.py @@ -0,0 +1,110 @@ +"""Provide tests for pattern.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest + +from junifer.datagrabber.pattern import PatternDataGrabber + + +def test_PatternDataGrabber() -> None: + """Test PatternDataGrabber.""" + + with pytest.raises(TypeError, match=r"`types` must be a list"): + PatternDataGrabber( + datadir="/tmp", + types="wrong", + patterns={"wrong": "pattern"}, + replacements="subject", + ) + + with pytest.raises(TypeError, match=r"`types` must be a list of strings"): + PatternDataGrabber( + datadir="/tmp", + types=[1, 2, 3], + patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, + replacements="subject", + ) + + with pytest.raises(ValueError, match=r"must have the same length"): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, + replacements=1, + ) + + with pytest.raises(TypeError, match=r"`patterns` must be a dict"): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns="wrong", + replacements="subject", + ) + + with pytest.raises( + ValueError, match=r"`patterns` must have the same length" + ): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"wrong": "pattern"}, + replacements="subject", + ) + + with pytest.raises( + ValueError, match=r"`patterns` must contain all `types`" + ): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"wrong": "pattern", "func": "pattern"}, + replacements="subject", + ) + + with pytest.raises(TypeError, match=r"must be a list of strings"): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"func": "func/test", "anat": "anat/test"}, + replacements=1, + ) + + with pytest.warns(RuntimeWarning, match=r"not part of any pattern"): + PatternDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={ + "func": "func/{subject}.nii", + "anat": "anat/{subject}.nii", + }, + replacements=["subject", "wrong"], + ) + + datagrabber = PatternDataGrabber( + datadir="/tmp/data", + types=["func", "anat"], + patterns={"func": "func/{subject}.nii", "anat": "anat/{subject}.nii"}, + replacements="subject", + ) + assert datagrabber.datadir == Path("/tmp/data") + assert datagrabber.types == ["func", "anat"] + assert datagrabber.replacements == ["subject"] + + datagrabber = PatternDataGrabber( + datadir=Path("/tmp/data"), + types=["func", "anat"], + patterns={ + "func": "func/{subject}.nii", + "anat": "anat/{subject}_{session}.nii", + }, + replacements=["subject", "session"], + ) + assert datagrabber.datadir == Path("/tmp/data") + assert datagrabber.types == ["func", "anat"] + assert datagrabber.replacements == ["subject", "session"] diff --git a/junifer/datagrabber/tests/test_pattern_datalad.py b/junifer/datagrabber/tests/test_pattern_datalad.py new file mode 100644 index 000000000..105918c44 --- /dev/null +++ b/junifer/datagrabber/tests/test_pattern_datalad.py @@ -0,0 +1,178 @@ +"""Provide tests for pattern_datalad.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest + +from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber + + +_testing_dataset = { + "example_bids": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids", + "id": "e2ce149bd723088769a86c72e57eded009258c6b", + }, + "example_bids_ses": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "id": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + }, +} + + +def test_bids_pattern_datalad_datagrabber_missing_uri() -> None: + """Test check of missing URI in pattern datalad datagrabber.""" + with pytest.raises(ValueError, match=r"`uri` must be provided"): + PatternDataladDataGrabber( + datadir=None, + types=[], + patterns={}, + replacements=[], + ) + + +def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None: + """Test a subject-based BIDS datalad datagrabber. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Define types + types = ["T1w", "bold"] + # Define patterns + patterns = { + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", + } + # Define replacements + replacements = ["subject"] + + repo_uri = _testing_dataset["example_bids"]["uri"] + rootdir = "example_bids" + repo_commit = _testing_dataset["example_bids"]["id"] + + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [f"sub-{i:02d}" for i in range(1, 10)] + assert set(subs) == set(expected_subs) + + for elem in dg: + t_sub = dg[elem] + assert "path" in t_sub["T1w"] + assert t_sub["T1w"]["path"] == ( + dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" + ) + assert "path" in t_sub["bold"] + assert t_sub["bold"]["path"] == ( + dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" + ) + + assert "meta" in t_sub + assert "datagrabber" in t_sub["meta"] + dg_meta = t_sub["meta"]["datagrabber"] + assert "class" in dg_meta + assert dg_meta["class"] == "PatternDataladDataGrabber" + assert "uri" in dg_meta + assert dg_meta["uri"] == repo_uri + assert "dataset_commit_id" in dg_meta + assert dg_meta["dataset_commit_id"] == repo_commit + + with open(t_sub["T1w"]["path"], "r") as f: + assert f.readlines()[0] == "placeholder" + + # datadir = tmp_path / "dataset" # Need this for testing + # patterns = { + # "T1w": "{subject}/anat/{subject}_T*w.nii.gz", + # "bold": "{subject}/func/{subject}_task-rest_*.nii.gz", + # } + # with PatternDataladDataGrabber( + # rootdir=rootdir, + # uri=repo_uri, + # types=types, + # patterns=patterns, + # datadir=datadir, + # replacements=replacements, + # ) as dg: + # assert dg.datadir == datadir / rootdir + # for elem in dg: + # t_sub = dg[elem] + # assert "path" in t_sub["T1w"] + # assert t_sub["T1w"]["path"] == ( + # dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" + # ) + # assert "path" in t_sub["bold"] + # assert t_sub["bold"]["path"] == ( + # dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" + # ) + + +def test_bids_PatternDataladDataGrabber_session(): + """Test a subject and session-based BIDS datalad datagrabber.""" + types = ["T1w", "bold"] + patterns = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + replacements = ["subject", "session"] + + with pytest.raises(ValueError, match=r"`uri` must be provided"): + PatternDataladDataGrabber( + datadir=None, + types=types, + patterns=patterns, + replacements=replacements, + ) + + repo_uri = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + # repo_commit = _testing_dataset['example_bids_ses']['id'] + + # With T1W and bold, only 2 sessions are available + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 3) + for i in range(1, 10) + ] + assert set(subs) == set(expected_subs) + + # Test with a different T1w only, it should have 3 sessions + types = ["T1w"] + patterns = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 4) + for i in range(1, 10) + ] + assert set(subs) == set(expected_subs) diff --git a/junifer/datagrabber/utils.py b/junifer/datagrabber/utils.py new file mode 100644 index 000000000..d30c11c13 --- /dev/null +++ b/junifer/datagrabber/utils.py @@ -0,0 +1,77 @@ +"""Provide utility functions for the datagrabber sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..utils import raise_error, warn_with_log + + +def validate_types(types: List[str]) -> None: + """Validate the types. + + Parameters + ---------- + types : list of str + The object to validate. + + """ + if not isinstance(types, list): + raise_error(msg="`types` must be a list", klass=TypeError) + if any(not isinstance(x, str) for x in types): + raise_error(msg="`types` must be a list of strings", klass=TypeError) + + +def validate_replacements( + replacements: List[str], patterns: Dict[str, str] +) -> None: + """Validate the replacements. + + Parameters + ---------- + replacements : list of str + The object to validate. + patterns : dict + The patterns to validate against. + + """ + if not isinstance(replacements, list): + raise_error(msg="`replacements` must be a list.", klass=TypeError) + if any(not isinstance(x, str) for x in replacements): + raise_error( + msg="`replacements` must be a list of strings.", klass=TypeError + ) + + for x in replacements: + if all(x not in y for y in patterns.values()): + warn_with_log(msg=f"Replacement {x} is not part of any pattern.") + + +def validate_patterns(types: List[str], patterns: Dict[str, str]) -> None: + """Validate the patterns. + + Parameters + ---------- + types : list of str + The types list. + patterns : dict + The object to validate. + + """ + # Validate the types + validate_types(types) + if not isinstance(patterns, dict): + raise_error(msg="`patterns` must be a dict.", klass=TypeError) + # Unequal length of objects + if len(types) != len(patterns): + raise_error( + msg="`types` and `patterns` must have the same length.", + klass=ValueError, + ) + + if any(x not in patterns for x in types): + raise_error( + msg="`patterns` must contain all `types`", klass=ValueError + ) diff --git a/junifer/datareader/__init__.py b/junifer/datareader/__init__.py index 4f1b6e6ca..19212063c 100644 --- a/junifer/datareader/__init__.py +++ b/junifer/datareader/__init__.py @@ -1,5 +1,8 @@ +"""Provide imports for datareader sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from .default import DefaultDataReader \ No newline at end of file +from .default import DefaultDataReader diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index d06494d63..dd3877d2f 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -1,66 +1,118 @@ +"""Provide class for default data reader.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL from pathlib import Path +from typing import Dict, List import nibabel as nib import pandas as pd -from ..utils.logging import logger -from ..markers.base import PipelineStepMixin +from ..pipeline.pipeline_mixin import PipelineStepMixin +from ..utils.logging import logger, warn_with_log -# Map each filenanm end to a kind + +# Map each file extension to a kind _extensions = { - '.nii': 'NIFTI', - '.nii.gz': 'NIFTI', - '.csv': 'CSV', - '.tsv': 'TSV' - + ".nii": "NIFTI", + ".nii.gz": "NIFTI", + ".csv": "CSV", + ".tsv": "TSV", } # Map each kind to a function and arguments _readers = {} -_readers['NIFTI'] = dict(func=nib.load, params=None) -_readers['CSV'] = dict(func=pd.read_csv, params=None) -_readers['TSV'] = dict(func=pd.read_csv, params={'sep': '\t'}) +_readers["NIFTI"] = {"func": nib.load, "params": None} +_readers["CSV"] = {"func": pd.read_csv, "params": None} +_readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}} class DefaultDataReader(PipelineStepMixin): + """Mixin class for default data reader.""" - def validate_input(self, input): + # TODO: complete type annotations + def validate_input(self, input: List[str]) -> None: + """Validate input. + + Parameters + ---------- + input + + """ # Nothing to validate, any input is fine pass + # TODO: complete type annotations def get_output_kind(self, input): + """Get output kind. + + Parameters + ---------- + input + + """ # It will output the same kind of data as the input return input - def fit_transform(self, input, params=None): + # TODO: complete type annotations + def fit_transform(self, input, params=None) -> Dict: + """Fit and transform. + + Parameters + ---------- + input + params + + Returns + ------- + dict + + """ # For each kind of data, try to read it - out = {} + + # out is the same, but with the 'data' key set in + # each kind dictionary, except for meta + out = input.copy() if params is None: params = {} for kind in input.keys(): - t_path = input[kind] + if kind == "meta": + out["meta"] = input["meta"] + continue + if "path" not in input[kind]: + warn_with_log( + f"Input kind {kind} does not provide a path. Skipping." + ) + continue + t_path = input[kind]["path"] t_params = params.get(kind, {}) + + # Convert to Path if datareader is not well done if not isinstance(t_path, Path): t_path = Path(t_path) - out[kind] = {'path': t_path} + out[kind]["path"] = t_path + logger.info(f"Reading {kind} from {t_path.as_posix()}") fread = None fname = t_path.name.lower() for ext, ftype in _extensions.items(): if fname.endswith(ext): - logger.info(f'Reading {ftype} file {t_path.as_posix()}') - reader_func = _readers[ftype]['func'] - reader_params = _readers[ftype]['params'] + logger.info(f"{kind} is type {ftype}") + reader_func = _readers[ftype]["func"] + reader_params = _readers[ftype]["params"] if reader_params is not None: t_params.update(reader_params) - logger.debug(f'Calling {reader_func} with {t_params}') + logger.debug(f"Calling {reader_func} with {t_params}") fread = reader_func(t_path, **t_params) break if fread is None: logger.info( - f'Unknown file type {t_path.as_posix()}, skipping reading') - out[kind]['data'] = fread + f"Unknown file type {t_path.as_posix()}, skipping reading" + ) + out[kind]["data"] = fread + if "meta" not in out: + out["meta"] = {} + out["meta"]["datareader"] = self.get_meta() return out diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index e1781d643..468583f7f 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -1,128 +1,168 @@ +"""Provide tests for default data reader.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL from pathlib import Path -import tempfile -from numpy.testing import assert_array_equal + import nibabel as nib -from nibabel import testing as nib_testing import pandas as pd +import pytest +from nibabel import testing as nib_testing +from numpy.testing import assert_array_equal from pandas.testing import assert_frame_equal from junifer.datareader import DefaultDataReader -def test_validation(): - """Test validating input/output""" - kinds = [ - ['T1w', 'BOLD', 'T2', 'dwi'], - [], - None, - ['whatever'] - ] +@pytest.mark.parametrize( + "kind", [["T1w", "BOLD", "T2", "dwi"], [], None, ["whatever"]] +) +def test_validation(kind) -> None: + """Test validating input/output. + Parameters + ---------- + kind : list of str or str or None + The parametrized kind of data. + + """ reader = DefaultDataReader() - - for t_kind in kinds: - assert reader.validate_input(t_kind) is None - assert reader.get_output_kind(t_kind) == t_kind - assert reader.validate(t_kind) == t_kind + assert reader.validate_input(kind) is None + assert reader.get_output_kind(kind) == kind + assert reader.validate(kind) == kind -def test_read_nifti(): - """Test reading NIFTI files""" +def test_meta() -> None: + """Test reader metadata.""" + reader = DefaultDataReader() + t_meta = reader.get_meta() + assert t_meta["class"] == "DefaultDataReader" + + nib_data_path = Path(nib_testing.data_path) + t_path = nib_data_path / "example4d.nii.gz" + input = {"bold": {"path": t_path}} + output = reader.fit_transform(input) + assert "meta" in output + assert "datareader" in output["meta"] + assert "class" in output["meta"]["datareader"] + assert output["meta"]["datareader"]["class"] == "DefaultDataReader" + + +@pytest.mark.parametrize( + "fname", ["example4d.nii.gz", "reoriented_anat_moved.nii"] +) +def test_read_nifti(fname: str) -> None: + """Test reading NIFTI files. + + Parameters + ---------- + fname : str + The parametrized NIfTI file names for testing. + + """ reader = DefaultDataReader() nib_data_path = Path(nib_testing.data_path) - for fname in ['example4d.nii.gz', - 'reoriented_anat_moved.nii']: - t_path = nib_data_path / fname + t_path = nib_data_path / fname - input = {'bold': t_path} - output = reader.fit_transform(input) - - assert isinstance(output, dict) - assert 'bold' in output - assert isinstance(output['bold'], dict) - assert 'path' in output['bold'] - assert 'data' in output['bold'] - - read_img = output['bold']['data'] - - t_read_img = nib.load(t_path) - assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) - - -def test_read_unknown(): - """Test (not) reading unknown files""" - reader = DefaultDataReader() - nib_data_path = Path(nib_testing.data_path) - - anat_path = nib_data_path / 'reoriented_anat_moved.nii' - whatever_path = nib_data_path / 'unexistant.unkwnownextension' - - input = {'anat': anat_path, 'whatever': whatever_path} + input = {"bold": {"path": t_path}} output = reader.fit_transform(input) assert isinstance(output, dict) - assert 'anat' in output - assert isinstance(output['anat'], dict) - assert 'path' in output['anat'] - assert isinstance(output['anat']['path'], Path) - assert 'data' in output['anat'] - assert output['anat']['data'] is not None + assert "bold" in output + assert isinstance(output["bold"], dict) + assert "path" in output["bold"] + assert "data" in output["bold"] - assert isinstance(output['whatever'], dict) - assert 'path' in output['whatever'] - assert isinstance(output['whatever']['path'], Path) - assert 'data' in output['whatever'] - assert output['whatever']['data'] is None + read_img = output["bold"]["data"] + + t_read_img = nib.load(t_path) + assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) + + input = {"bold": {"path": t_path.as_posix()}} + output2 = reader.fit_transform(input) + assert output["bold"]["path"] == output2["bold"]["path"] -def test_read_csv(): - """Test reading CSV files""" - d = {'col1': [1, 2, 3, 4, 5], 'col2': [3, 4, 5, 6, 7]} +def test_read_unknown() -> None: + """Test (not) reading unknown files.""" + reader = DefaultDataReader() + nib_data_path = Path(nib_testing.data_path) + + anat_path = nib_data_path / "reoriented_anat_moved.nii" + whatever_path = nib_data_path / "unexistent.unkwnownextension" + + input = {"anat": {"path": anat_path}, "whatever": {"path": whatever_path}} + output = reader.fit_transform(input) + + assert isinstance(output, dict) + assert "anat" in output + assert isinstance(output["anat"], dict) + assert "path" in output["anat"] + assert isinstance(output["anat"]["path"], Path) + assert "data" in output["anat"] + assert output["anat"]["data"] is not None + + assert isinstance(output["whatever"], dict) + assert "path" in output["whatever"] + assert isinstance(output["whatever"]["path"], Path) + assert "data" in output["whatever"] + assert output["whatever"]["data"] is None + + +def test_read_csv(tmp_path: Path) -> None: + """Test reading CSV files. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + d = {"col1": [1, 2, 3, 4, 5], "col2": [3, 4, 5, 6, 7]} df = pd.DataFrame(d) - with tempfile.TemporaryDirectory() as tmpdir: - tmpdir = Path(tmpdir) - df.to_csv(tmpdir / 'test.csv') - reader = DefaultDataReader() - input = {'csv': tmpdir / 'test.csv'} - output = reader.fit_transform(input) + df.to_csv(tmp_path / "test_read_csv.csv") - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + reader = DefaultDataReader() + input = {"csv": {"path": tmp_path / "test_read_csv.csv"}} + output = reader.fit_transform(input) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] - df.to_csv(tmpdir / 'test.csv', sep=';') - input = {'csv': tmpdir / 'test.csv'} - params = {'csv': {'sep': ';'}} - output = reader.fit_transform(input, params) + read_df = output["csv"]["data"][["col1", "col2"]] + assert_frame_equal(df, read_df) - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + df.to_csv(tmp_path / "test_read_csv.csv", sep=";") + input = {"csv": {"path": tmp_path / "test_read_csv.csv"}} + params = {"csv": {"sep": ";"}} + output = reader.fit_transform(input, params) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] - df.to_csv(tmpdir / 'test.tsv', sep='\t') - input = {'csv': tmpdir / 'test.tsv'} - output = reader.fit_transform(input) + read_df = output["csv"]["data"][["col1", "col2"]] + assert_frame_equal(df, read_df) - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + df.to_csv(tmp_path / "test_read_csv.tsv", sep="\t") + input = {"csv": {"path": tmp_path / "test_read_csv.tsv"}} + output = reader.fit_transform(input) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] + + read_df = output["csv"]["data"][["col1", "col2"]] + # Check if dataframes are equal + assert_frame_equal(df, read_df) diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 4e235bc8b..c73e6f8d5 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -1,3 +1,9 @@ +"""Provide imports for markers sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse -# License: AGPL \ No newline at end of file +# License: AGPL + +from .base import BaseMarker +from .collection import MarkerCollection +from .parcel import ParcelAggregation diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 1aeb3b284..d08348868 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -1,62 +1,155 @@ +"""Provide base class for markers.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -class PipelineStepMixin(): +from typing import Dict, List, Optional, Union - @property - def name(self): - return self.__class__.__name__ +from ..pipeline.pipeline_mixin import PipelineStepMixin +from ..utils import logger, raise_error - def validate_input(self, input): - """Validate the input to the pipeline step. + +class BaseMarker(PipelineStepMixin): + """Base class for all markers. + + Parameters + ---------- + on : list of str + The kind of data to work on. + name : str, optional + The name of the marker (default None). + + """ + + def __init__( + self, on: Union[List[str], str], name: Optional[str] = None + ) -> None: + """Initialize the class.""" + if not isinstance(on, list): + on = [on] + self._valid_inputs = on + self.name = self.__class__.__name__ if name is None else name + + def get_meta(self, kind: str) -> Dict: + """Get metadata. Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. - - Raises - ------ - ValueError: - If the input does not have the required data. - """ - raise NotImplementedError('validate_input not implemented') - - def get_output_kind(self, input): - """Get the kind of the pipeline step. - - Parameters - ---------- - input : Junifer Data dictionary - The input to the pipeline step. + kind : str + The kind of pipeline step. Returns ------- - output : Junifer Data dictionary - The output of the pipeline step. - """ - raise NotImplementedError('get_output_kind not implemented') + dict + The metadata as a dictionary. - def validate(self, input): - """Validate the the pipeline step. + """ + s_meta = super().get_meta() + # same marker can be "fit"ted into different kinds, so the name + # is created from the kind and the name of the marker + s_meta["name"] = f"{kind}_{self.name}" + s_meta["kind"] = kind + return {"marker": s_meta} + + def validate_input(self, input: List[str]) -> None: + """Validate input. Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. - - Returns - ------- - output : Junifer Data dictionary - The output of the pipeline step. + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. Raises ------ - ValueError: + ValueError If the input does not have the required data. - """ - self.validate_input(input) - return self.get_output_kind(input) - def fit_transform(self, input): - raise NotImplementedError('fit_transform not implemented') + """ + if not any(x in input for x in self._valid_inputs): + raise_error( + "Input does not have the required data." + f"\t Input: {input}" + f"\t Required (any of): {self._valid_inputs}" + ) + + def get_output_kind(self, input: List[str]) -> List[str]: + """Get output kind. + + Parameters + ---------- + input : list of str + The input to the marker. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + list of str + The updated list of output kinds, as storage possibilities. + + """ + raise_error(msg="compute() not implemented", klass=NotImplementedError) + + def compute(self, input: Dict) -> Dict: + """Compute. + + Parameters + ---------- + input : Dict[str, Dict] + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + dict + The computed result as dictionary. + + """ + raise_error(msg="compute() not implemented", klass=NotImplementedError) + + # TODO: complete type annotations + def store(self, kind: str, out: Dict, storage) -> None: + """Store. + + Parameters + ---------- + input + out + + """ + raise_error(msg="store() not implemented", klass=NotImplementedError) + + # TODO: complete type annotations + def fit_transform(self, input: Dict[str, Dict], storage=None) -> Dict: + """Fit and transform. + + Parameters + ---------- + input + storage + + Returns + ------- + dict + + """ + out = {} + meta = input.get("meta", {}) + for kind in self._valid_inputs: + if kind in input.keys(): + logger.info(f"Computing {kind}") + t_input = input[kind] + t_meta = meta.copy() + t_meta.update(t_input.get("meta", {})) + t_meta.update(self.get_meta(kind)) + t_out = self.compute(t_input) + t_out.update(meta=t_meta) + if storage is not None: + logger.info(f"Storing in {storage}") + self.store(kind, t_out, storage) + else: + logger.info("No storage specified, returning dictionary") + out[kind] = t_out + + return out diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 9d06aa0a5..1897783f7 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -1,70 +1,111 @@ +"""Provide class for marker collection.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL +from collections import Counter +from typing import Dict, List, Optional + +from ..datareader.default import DefaultDataReader +from ..markers.base import BaseMarker +from ..pipeline import PipelineStepMixin +from ..storage.base import BaseFeatureStorage from ..utils import logger -from ..datareader import DefaultDataReader -class MarkerCollection(): - def __init__(self, markers, datareader=None, preprocessing=None, - storage=None): +class MarkerCollection: + """Class for marker collection. + + Parameters + ---------- + markers + datareader + preprocessing + storage + + """ + + def __init__( + self, + markers: List[BaseMarker], + datareader: Optional[PipelineStepMixin] = None, + preprocessing: Optional[PipelineStepMixin] = None, + storage: Optional[BaseFeatureStorage] = None, + ): + """Initialize the class.""" + # Check that the markers have different names + marker_names = [m.name for m in markers] + if len(set(marker_names)) != len(marker_names): + counts = Counter(marker_names) + raise ValueError( + "Markers must have different names. " + f"Current names are: {counts}" + ) + self._markers = markers if datareader is None: datareader = DefaultDataReader() self._datareader = datareader self._preprocessing = preprocessing - self._markers = markers self._storage = storage - def fit(self, input): + def fit(self, input: Dict[str, Dict]) -> Optional[Dict]: """Fit the pipeline. Parameters ---------- - input : Junifer Data dictionary (input) + input The input data to fit the pipeline on. Should be the output of indexing the DataGrabber with one element. Returns ------- - output : dict[str -> object] + output : dict or None The output of the pipeline. Each key represents a marker name and the values are the computer marker values. If the pipeline has a storage configured, then the output will be None. + """ - logger.info('Fitting pipeline') + logger.info("Fitting pipeline") data = self._datareader.fit_transform(input) if self._preprocessing is not None: - logger.info('Preprocessing data') + logger.info("Preprocessing data") data = self._preprocessing.fit_transform(data) out = {} for marker in self._markers: - logger.info(f'Fitting marker {marker.name}') + logger.info(f"Fitting marker {marker.name}") m_value = marker.fit_transform(data, storage=self._storage) if self._storage is None: out[marker.name] = m_value - + logger.info("Marker collection fitting done") return None if self._storage else out - def validate(self, datagrabber): + # TODO: complete type annotations + def validate(self, datagrabber) -> None: """Validate the pipeline. Without doing any computation, check if the Marker Collection can be fit without problems. That is, the data required for each marker is present and streamed down the steps. Also, if a storage is configured, check that the storage can handle the markers output. - """ - logger.info('Validating Marker Collection') - t_data = datagrabber.get_output_kind() - logger.info(f'DataGrabber output type: {t_data}') - logger.info(f'Validating Data Reader:') + Parameters + ---------- + datagrabber + + """ + logger.info("Validating Marker Collection") + t_data = datagrabber.get_types() + logger.info(f"DataGrabber output type: {t_data}") + + logger.info("Validating Data Reader:") t_data = self._datareader.validate(t_data) - logger.info(f'Data Reader output type: {t_data}') + logger.info(f"Data Reader output type: {t_data}") for marker in self._markers: - logger.info(f'Validating Marker: {marker.name}') + logger.info(f"Validating Marker: {marker.name}") m_data = marker.validate(t_data) - logger.info(f'Marker output type: {m_data}') + logger.info(f"Marker output type: {m_data}") if self._storage is not None: - logger.info(f'Validating storage for {marker.name}') + logger.info(f"Validating storage for {marker.name}") self._storage.validate(m_data) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py new file mode 100644 index 000000000..386ad8dbb --- /dev/null +++ b/junifer/markers/parcel.py @@ -0,0 +1,144 @@ +"""Provide class for parcel aggregation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +import numpy as np +from nilearn.image import math_img, resample_to_img +from nilearn.maskers import NiftiMasker + +from ..api.decorators import register_marker +from ..data import load_atlas +from ..stats import get_aggfunc_by_name +from ..utils import logger +from .base import BaseMarker + + +@register_marker +class ParcelAggregation(BaseMarker): + """Class for parcel aggregation. + + Parameters + ---------- + atlas + method + method_params + on + name + + """ + + def __init__( + self, atlas, method, method_params=None, on=None, name=None + ) -> None: + """Initialize the class.""" + self.atlas = atlas + self.method = method + self.method_params = {} if method_params is None else method_params + if on is None: + on = ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] + super().__init__(on=on, name=name) + + def get_output_kind(self, input: List[str]) -> List[str]: + """Get output kind. + + Parameters + ---------- + input : list of str + The kind of data to work on. + + Returns + ------- + str + The kind of output. + + """ + outputs = [] + for t_input in input: + if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + outputs.append("table") + elif input in ["BOLD"]: + outputs.append("timeseries") + else: + raise ValueError(f"Unknown input kind for {t_input}") + return outputs + + # TODO: complete type annotations + def store(self, kind: str, out, storage) -> None: + """Store. + + Parameters + ---------- + kind + out + storage + + """ + logger.debug(f"Storing {kind} in {storage}") + if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + storage.store_table(**out) + if kind in ["BOLD"]: + storage.store_timeseries(**out) + + # TODO: complete type annotations + def compute(self, input) -> Dict: + """Compute. + + Parameters + ---------- + input + + Returns + ------- + dict + The computed result as dictionary. + + """ + t_input = input["data"] + logger.debug(f"Parcel aggregation using {self.method}") + agg_func = get_aggfunc_by_name( + self.method, func_params=self.method_params + ) + # Get the min of the voxels sizes and use it as the resolution + resolution = np.min(t_input.header.get_zooms()[:3]) # type: ignore + t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution) + atlas_img_res = resample_to_img( + t_atlas, + t_input, + interpolation="nearest", + ) + atlas_bin = math_img( + "img != 0", + img=atlas_img_res, + ) + logger.debug("Masking") + masker = NiftiMasker( + atlas_bin, target_affine=t_input.affine + ) # type: ignore + + # Mask the input data and the atlas + data = masker.fit_transform(t_input) + atlas_values = masker.transform(atlas_img_res) + atlas_values = np.squeeze(atlas_values).astype(int) + + # Get the values for each parcel and apply agg function + logger.debug("Computing ROI means") + atlas_roi_vals = sorted(np.unique(atlas_values)) + out_labels = [] + out_values = [] + # Iterate over the parcels (existing) + for t_v in atlas_roi_vals: + t_values = agg_func(data[:, atlas_values == t_v], axis=-1) + out_values.append(t_values) + # Update the labels just in case a parcel has no voxels + # in it + out_labels.append(t_labels[t_v - 1]) + + out_values = np.array(out_values).T + out = {"data": out_values, "columns": out_labels} + if out_values.shape[0] > 1: + out["row_names"] = "scan" + return out diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py new file mode 100644 index 000000000..8ef21777d --- /dev/null +++ b/junifer/markers/tests/test_collection.py @@ -0,0 +1,165 @@ +"""Provide tests for marker collection.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import pytest +from numpy.testing import assert_array_equal + +from junifer.datareader.default import DefaultDataReader +from junifer.markers import MarkerCollection, ParcelAggregation +from junifer.pipeline import PipelineStepMixin +from junifer.storage import SQLiteFeatureStorage +from junifer.testing.datagrabbers import OasisVBMTestingDatagrabber + + +def test_marker_collection_incorrect_markers() -> None: + """Test incorrect markers for MarkerCollection.""" + wrong_markers = [ + ParcelAggregation( + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), + ParcelAggregation( + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), + ] + with pytest.raises(ValueError, match=r"must have different names"): + MarkerCollection(wrong_markers) + + +def test_marker_collection(): + """Test MarkerCollection.""" + markers = [ + ParcelAggregation( + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), + ParcelAggregation( + atlas="Schaefer100x7", method="std", name="gmd_schaefer100x7_std" + ), + ParcelAggregation( + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + name="gmd_schaefer100x7_trim_mean90", + ), + ] + mc = MarkerCollection(markers=markers) + assert mc._markers == markers + assert mc._preprocessing is None + assert mc._storage is None + assert isinstance(mc._datareader, DefaultDataReader) + + # Create testing datagrabber + dg = OasisVBMTestingDatagrabber() + mc.validate(dg) + + with dg: + input = dg["sub-01"] + out = mc.fit(input) + assert out is not None + assert isinstance(out, dict) + assert len(out) == 3 + assert "gmd_schaefer100x7_mean" in out + assert "gmd_schaefer100x7_std" in out + assert "gmd_schaefer100x7_trim_mean90" in out + + for t_marker in markers: + t_name = t_marker.name + assert "VBM_GM" in out[t_name] + t_vbm = out[t_name]["VBM_GM"] + assert "data" in t_vbm + assert "columns" in t_vbm + assert "meta" in t_vbm + + # Test preprocessing + class BypassPreprocessing(PipelineStepMixin): + def fit_transform(self, input): + return input + + mc2 = MarkerCollection( + markers=markers, + preprocessing=BypassPreprocessing(), + datareader=DefaultDataReader(), + ) + assert isinstance(mc2._datareader, DefaultDataReader) + with dg: + input = dg["sub-01"] + out2 = mc2.fit(input) + assert out2 is not None + for t_marker in markers: + t_name = t_marker.name + assert_array_equal( + out[t_name]["VBM_GM"]["data"], out2[t_name]["VBM_GM"]["data"] + ) + + +def test_MarkerCollection_storage(tmp_path) -> None: + """Test marker collection with storage. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + markers = [ + ParcelAggregation( + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), + ParcelAggregation( + atlas="Schaefer100x7", method="std", name="gmd_schaefer100x7_std" + ), + ParcelAggregation( + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + name="gmd_schaefer100x7_trim_mean90", + ), + ] + # Test storage + dg = OasisVBMTestingDatagrabber() + + uri = tmp_path / "test_marker_collection_storage.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + mc = MarkerCollection( + markers=markers, storage=storage, datareader=DefaultDataReader() + ) + mc.validate(dg) + assert mc._storage is not None + assert mc._storage.uri == storage.uri + with dg: + input = dg["sub-01"] + out = mc.fit(input) + assert out is None + + mc2 = MarkerCollection(markers=markers, datareader=DefaultDataReader()) + mc2.validate(dg) + assert mc2._storage is None + + with dg: + input = dg["sub-01"] + out = mc2.fit(input) + + features = storage.list_features() + assert len(features) == 3 + feature_md5 = list(features.keys())[0] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_mean" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore + + feature_md5 = list(features.keys())[1] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_std" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore + + feature_md5 = list(features.keys())[2] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_trim_mean90" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore diff --git a/junifer/markers/tests/test_markers_base.py b/junifer/markers/tests/test_markers_base.py new file mode 100644 index 000000000..c6459d625 --- /dev/null +++ b/junifer/markers/tests/test_markers_base.py @@ -0,0 +1,78 @@ +"""Provide tests for base marker.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import List, Optional + +import pytest + +from junifer.markers.base import BaseMarker + + +@pytest.mark.parametrize( + "on, name, kind, expected_class, expected_name", + [ + (["bold", "dwi"], None, "bold", "BaseMarker", "bold_BaseMarker"), + (["bold", "dwi"], "mymarker", "dwi", "BaseMarker", "dwi_mymarker"), + ], +) +def test_base_marker_meta( + on: List[str], + name: Optional[str], + kind: str, + expected_class: str, + expected_name: str, +) -> None: + """Test metadata for BaseMarker. + + Parameters + ---------- + on : list of str + The parametrized kind of data to work on. + name : str or None + The parametrized name of the marker. + kind : str + The parametrized kind of data to get metadata for. + expected_class : str + The paramtrized expected class of the marker. + expected_name : str + The parametrized expected name of the marker. + + """ + base = BaseMarker(on=on, name=name) + t_meta = base.get_meta(kind=kind) + assert t_meta["marker"]["class"] == expected_class + assert t_meta["marker"]["name"] == expected_name + + +def test_BaseMarker() -> None: + """Test base class.""" + base = BaseMarker(on=["bold", "dwi"], name="mymarker") + input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}} + base.validate_input(list(input_.keys())) + + wrong_input = {"t2": {"path": "test"}} + with pytest.raises(ValueError): + base.validate_input(list(wrong_input.keys())) + + with pytest.raises(NotImplementedError): + base.get_output_kind(list(wrong_input.keys())) + + with pytest.raises(NotImplementedError): + base.fit_transform(input_) + + base.compute = lambda x: {"data": 1} # type: ignore + + out = base.fit_transform(input_) + assert out["bold"]["data"] == 1 + assert out["bold"]["meta"]["marker"]["name"] == "bold_mymarker" + assert out["bold"]["meta"]["marker"]["class"] == "BaseMarker" + + base2 = BaseMarker(on="bold", name="mymarker") + base2.compute = lambda x: {"data": 1} # type: ignore + out2 = base2.fit_transform(input_) + assert out2["bold"]["data"] == 1 + assert out2["bold"]["meta"]["marker"]["name"] == "bold_mymarker" + assert out2["bold"]["meta"]["marker"]["class"] == "BaseMarker" diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py new file mode 100644 index 000000000..1bd78a50f --- /dev/null +++ b/junifer/markers/tests/test_parcel.py @@ -0,0 +1,166 @@ +"""Provide test for parcel aggregation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import nibabel as nib +import numpy as np +from nilearn import datasets +from nilearn.image import concat_imgs, math_img, resample_to_img +from nilearn.maskers import NiftiLabelsMasker, NiftiMasker +from numpy.testing import assert_array_almost_equal, assert_array_equal +from scipy.stats import trim_mean + +from junifer.markers.parcel import ParcelAggregation + + +def test_ParcelAggregation_3D() -> None: + """Test ParcelAggregation object on 3D images.""" + # Get the testing atlas (for nilearn) + atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + + # Get the oasis VBM data + oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) + vbm = oasis_dataset.gray_matter_maps[0] + img = nib.load(vbm) + + # Mask atlas manually + atlas_img_res = resample_to_img( + atlas.maps, + img, + interpolation="nearest", + ) + atlas_bin = math_img( + "img != 0", + img=atlas_img_res, + ) + + # Create NiftiMasker + masker = NiftiMasker(atlas_bin, target_affine=img.affine) + data = masker.fit_transform(img) + atlas_values = masker.transform(atlas_img_res) + atlas_values = np.squeeze(atlas_values).astype(int) + + # Compute the mean manually + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = np.mean(data[:, atlas_values == t_v]) + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Create NiftiLabelsMasker + nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) + auto = nifti_masker.fit_transform(img) + + # Check that arrays are almost equal + assert_array_almost_equal(auto, manual) + + # Use the ParcelAggregation object + marker = ParcelAggregation( + atlas="Schaefer100x7", + method="mean", + name="gmd_schaefer100x7_mean", + on="VBM_GM", + ) # Test passing "on" as a keyword argument + input = dict(VBM_GM=dict(data=img)) + jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] + + assert jun_values3d_mean.ndim == 2 + assert jun_values3d_mean.shape[0] == 1 + assert_array_equal(manual, jun_values3d_mean) + + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {} + + # Test using another function (std) + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = np.std(data[:, atlas_values == t_v]) + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Use the ParcelAggregation object + marker = ParcelAggregation(atlas="Schaefer100x7", method="std") + input = dict(VBM_GM=dict(data=img)) + jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"] + + assert jun_values3d_std.ndim == 2 + assert jun_values3d_std.shape[0] == 1 + assert_array_equal(manual, jun_values3d_std) + + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "std" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {} + + # Test using another function with parameters + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = trim_mean( + data[:, atlas_values == t_v], proportiontocut=0.1, axis=None + ) # type: ignore + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Use the ParcelAggregation object + marker = ParcelAggregation( + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + ) + input = dict(VBM_GM=dict(data=img)) + jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"] + + assert jun_values3d_tm.ndim == 2 + assert jun_values3d_tm.shape[0] == 1 + assert_array_equal(manual, jun_values3d_tm) + + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "trim_mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {"proportiontocut": 0.1} + + +def test_ParcelAggregation_4D(): + """Test ParcelAggregation object on 4D images.""" + # Get the testing atlas (for nilearn) + atlas = datasets.fetch_atlas_schaefer_2018( + n_rois=100, yeo_networks=7, resolution_mm=2 + ) + + # Get the SPM auditory data: + subject_data = datasets.fetch_spm_auditory() + fmri_img = concat_imgs(subject_data.func) # type: ignore + + # Create NiftiLabelsMasker + nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) + auto4d = nifti_masker.fit_transform(fmri_img) + + # Create ParcelAggregation object + marker = ParcelAggregation(atlas="Schaefer100x7", method="mean") + input = dict(BOLD=dict(data=fmri_img)) + jun_values4d = marker.fit_transform(input)["BOLD"]["data"] + + assert jun_values4d.ndim == 2 + assert_array_equal(auto4d.shape, jun_values4d.shape) + assert_array_equal(auto4d, jun_values4d) + + meta = marker.get_meta("BOLD")["marker"] + assert meta["method"] == "mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "BOLD_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "BOLD" + assert meta["method_params"] == {} diff --git a/junifer/pipeline/__init__.py b/junifer/pipeline/__init__.py new file mode 100644 index 000000000..6d57d7b78 --- /dev/null +++ b/junifer/pipeline/__init__.py @@ -0,0 +1,6 @@ +"""Provide imports for pipeline sub-package.""" + +# Authors: Synchon Mandal +# License: AGPL + +from .pipeline_mixin import PipelineStepMixin diff --git a/junifer/pipeline/pipeline_mixin.py b/junifer/pipeline/pipeline_mixin.py new file mode 100644 index 000000000..915458810 --- /dev/null +++ b/junifer/pipeline/pipeline_mixin.py @@ -0,0 +1,107 @@ +"""Provide mixin class for pipeline step.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..utils import raise_error + + +class PipelineStepMixin: + """Mixin class for pipeline.""" + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as a dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith("_"): + t_meta[k] = v + return t_meta + + def validate_input(self, input: List[str]) -> None: + """Validate the input to the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + raise_error( + msg="Concrete classes need to implement validate_input().", + klass=NotImplementedError, + ) + + def get_output_kind(self, input: List[str]) -> List[str]: + """Get the kind of the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + list of str + The updated list of available Junifer Data dictionary keys after + the pipeline step. + + """ + raise_error( + msg="Concrete classes need to implement get_output_kind().", + klass=NotImplementedError, + ) + + def validate(self, input: List[str]) -> List[str]: + """Validate the the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. + + Returns + ------- + list of str + The output of the pipeline step. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + self.validate_input(input=input) + return self.get_output_kind(input=input) + + def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]: + """Fit and transform. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + """ + raise_error( + msg="Concrete classes need to implement fit_transform().", + klass=NotImplementedError, + ) diff --git a/junifer/pipeline/tests/test_pipeline_mixin.py b/junifer/pipeline/tests/test_pipeline_mixin.py new file mode 100644 index 000000000..afae53362 --- /dev/null +++ b/junifer/pipeline/tests/test_pipeline_mixin.py @@ -0,0 +1,27 @@ +"""Provide tests for pipeline mixin.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import pytest + +from junifer.pipeline.pipeline_mixin import PipelineStepMixin + + +def test_PipelineStepMixin() -> None: + """Test PipelineStepMixin.""" + mixin = PipelineStepMixin() + with pytest.raises(NotImplementedError): + mixin.validate_input([]) + with pytest.raises(NotImplementedError): + mixin.get_output_kind([]) + with pytest.raises(NotImplementedError): + mixin.fit_transform({}) + + +def test_pipeline_step_mixin_meta(): + """Test metadata for PipelineStepMixin.""" + pipemixin = PipelineStepMixin() + t_meta = pipemixin.get_meta() + assert t_meta["class"] == "PipelineStepMixin" diff --git a/junifer/preprocess/__init__.py b/junifer/preprocess/__init__.py index 4e235bc8b..b7016a59d 100644 --- a/junifer/preprocess/__init__.py +++ b/junifer/preprocess/__init__.py @@ -1,3 +1,7 @@ +"""Provide imports for preprocess sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse -# License: AGPL \ No newline at end of file +# License: AGPL + +from .confounds import BaseConfoundRemover diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py new file mode 100644 index 000000000..f61f29644 --- /dev/null +++ b/junifer/preprocess/confounds.py @@ -0,0 +1,487 @@ +"""Provide base class for confound removal.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from typing import TYPE_CHECKING, Dict, List, Optional, Union + +import numpy as np +import pandas as pd +from nilearn._utils.niimg_conversions import check_niimg_4d +from nilearn.image import clean_img +from nilearn.masking import compute_brain_mask + +from ..pipeline import PipelineStepMixin +from ..utils import logger, raise_error + + +if TYPE_CHECKING: + from nibabel import MGHImage, Nifti1Image, Nifti2Image + + +class BaseConfoundRemover(PipelineStepMixin): + """Base class for confound removal. + + Read confound files and select columns according to + a pre-defined strategy. + + Confound removal is based on `nilearn.image.clean_img`. + + Parameters + ---------- + strategy : dict, optional + The keys of the dictionary should correspond to names of noise + components to include: + - 'motion' + - 'wm_csf' + - 'global_signal' + The values of dictionary should correspond to types of confounds + extracted from each signal: + - 'basic': only the confounding time series + - 'power2': signal + quadratic term + - 'derivatives': signal + derivatives + - 'full': signal + deriv. + quadratic terms + power2 deriv. + (default None). + spike : float, optional + If None, no spike regressor is added. If spike is a float, it will + add a spike regressor for every point at which FD exceeds the + specified float (default None). + detrend : bool, Optional + If True, detrending will be applied on timeseries + (before confound removal) (default True). + standardize : bool, optional + If True, returned signals are set to unit variance (default True). + low_pass : float, optional + Low cutoff frequencies, in Hertz. If None, no filtering is applied + (default None). + high_pass : float, optional + High cutoff frequencies, in Hertz. If None, no filtering is + applied (default None). + t_r : float, optional + Repetition time, in second (sampling period). + If None, it will use t_r from nifti header (default None). + mask_img: Niimg-like object, optional + If provided, signal is only cleaned from voxels inside the mask. + If mask is provided, it should have same shape and affine as imgs. + If not provided, a mask is computed using + `nilearn.masking.compute_brain_mask` (default None). + + """ + + # lower priority + # TODO: implement more strategies from + # nilearn.interfaces.fmriprep.load_confounds for Felix's confound files, + # in particular scrubbing + # TODO: Implement read_confounds for fmriprep data + + def __init__( + self, + strategy: Optional[Dict[str, str]] = None, + spike: Optional[float] = None, + detrend: bool = True, + standardize: bool = True, + low_pass: Optional[float] = None, + high_pass: Optional[float] = None, + t_r: Optional[float] = None, + mask_img: Optional["Nifti1Image"] = None, + ) -> None: + """Initialise the class.""" + if strategy is None: + strategy = { + "motion": "full", + "wm_csf": "full", + "global_signal": "full", + } + self.strategy = strategy + self.spike = spike + self.detrend = detrend + self.standardize = standardize + self.low_pass = low_pass + self.high_pass = high_pass + self.t_r = t_r + self.mask_img = mask_img + + self._valid_components = ["motion", "wm_csf", "global_signal"] + self._valid_confounds = ["basic", "power2", "derivatives", "full"] + + if any(not isinstance(k, str) for k in strategy.keys()): + raise_error("Strategy keys must be strings", ValueError) + + if any(not isinstance(v, str) for v in strategy.values()): + raise_error("Strategy values must be strings", ValueError) + + if any(x not in self._valid_components for x in strategy.keys()): + raise_error( + msg=f"Invalid component names {list(strategy.keys())}. " + f"Valid components are {self._valid_components}.\n" + f"If any of them is a valid parameter in " + "nilearn.interfaces.fmriprep.load_confounds we may " + "include it in the future", + klass=ValueError, + ) + + if any(x not in self._valid_confounds for x in strategy.values()): + raise_error( + msg=f"Invalid component names {list(strategy.values())}. " + f"Valid confound types are {self._valid_confounds}.\n" + f"If any of them is a valid parameter in " + "nilearn.interfaces.fmriprep.load_confounds we may " + "include it in the future", + klass=ValueError, + ) + + def validate_input(self, input: List[str]) -> None: + """Validate the input to the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + _required_inputs = ["BOLD", "confounds"] + if any(x not in input for x in _required_inputs): + raise_error( + msg="Input does not have the required data. \n" + f"Input: {input} \n" + f"Required (all off): {_required_inputs} \n", + klass=ValueError, + ) + + def get_output_kind(self, input: List[str]) -> List[str]: + """Get the kind of the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + list of str + The updated list of available Junifer Data dictionary keys after + the pipeline step. + + """ + # Does not add any new keys + return input + + # TODO: complete type annotations + def _pick_confounds(self, input): + """Select relevant confounds from the specified file.""" + to_select = [] + confounds_df = input["data"] + confounds_spec = input["names"]["spec"] + # for every confound there is a derivative + # and for every confound + derivative there should be squares + derivatives_to_compute = input["names"].get("derivatives", {}) + squares_to_compute = input["names"].get("squares", {}) + spike_name = input["names"]["spike"] + + # Get all the column names according to the strategy + for comp, param in self.strategy.items(): + to_select.extend(confounds_spec[comp][param]) + + # Add derivatives if needed + to_compute = [x in derivatives_to_compute.keys() for x in to_select] + out_df = confounds_df.copy() + if any(to_compute): + for t_dst, t_src in derivatives_to_compute.items(): + out_df[t_dst] = np.append( # type: ignore + np.diff(out_df[t_src]), 0 + ) # type: ignore + + # Add squares (of base confounds and derivatives) if needed + to_compute = [x in squares_to_compute.keys() for x in to_select] + if any(to_compute): + for t_dst, t_src in squares_to_compute.items(): + out_df[t_dst] = out_df[t_src] ** 2 + out_df = out_df[to_select] + + # add binary spike regressor if needed at given threshold + if self.spike is not None: + fd = confounds_df[spike_name].copy() + fd.loc[fd > self.spike] = 1 + fd.loc[fd != 1] = 0 + out_df["spike"] = fd + + return out_df + + def _remove_confounds( + self, bold_img: "Nifti1Image", confounds_df: pd.DataFrame + ) -> Union["Nifti1Image", "Nifti2Image", "MGHImage", List]: + """Remove confounds from the BOLD data.""" + """Remove confounds from the BOLD image. + + Parameters + ---------- + bold_img : Niimg-like object + 4D image. The signals in the last dimension are filtered + (see http://nilearn.github.io/manipulating_images/input_output.html + for a detailed description of the valid input types). + confounds_df : pd.DataFrame + Dataframe containing confounds to remove. Number of rows should + correspond to number of volumes in the BOLD image. + + Returns + -------- + Niimg-like object + Input image with confounds removed. + + """ + confounds_array = confounds_df.values + + t_r = self.t_r + if t_r is None: + logger.info("No `t_r` specified, using t_r from nifti header") + zooms = bold_img.header.get_zooms() # type: ignore + t_r = zooms[3] + logger.info( + f"Read t_r from nifti header: {t_r}", + ) + + mask_img = self.mask_img + if mask_img is None: + logger.info("Computing brain mask from image") + mask_img = compute_brain_mask(bold_img) + + clean_bold = clean_img( + imgs=bold_img, + detrend=self.detrend, + standardize=self.standardize, + confounds=confounds_array, + low_pass=self.low_pass, + high_pass=self.high_pass, + t_r=t_r, + mask_img=mask_img, + ) + + return clean_bold + + # TODO: complete type annotations + def _validate_data(self, input): + """Validate input data.""" + # Bold must be 4D niimg + check_niimg_4d(input["BOLD"]["data"]) + + # Confounds must be a dataframe + if not isinstance(input["confounds"]["data"], pd.DataFrame): + raise_error( + "confounds data must be a pandas dataframe", ValueError + ) + + confound_df = input["confounds"]["data"] + bold_img = input["BOLD"]["data"] + if bold_img.get_fdata().shape[3] != len(confound_df): + raise_error( + "Image time series and confounds have different length!\n" + f"\tImage time series: { bold_img.get_fdata().shape[3]}\n" + f"\tConfounds: {len(confound_df)}" + ) + + # Check the column names of the dataframe and the spec + # spec must be a dictionary: + # { + # 'motion': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]}, + # 'wm_csf': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]} + # 'global_signal': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]} + # } + + # Check the columns in the dataframe + conf_spec = input["confounds"]["names"]["spec"] + if any(x not in conf_spec.keys() for x in self._valid_components): + raise_error( + "All of the component types must be in the confounds data " + "object `spec`. Please check your datagrabber.", + ValueError, + ) + + if any( + x not in v.keys() + for x in self._valid_confounds + for v in conf_spec.values() + ): + raise_error( + "All of the confound types must be in the confounds data " + "object `spec`. Please check your datagrabber.", + ValueError, + ) + + spike_name = input["confounds"]["names"]["spike"] + + derivatives_to_compute = input["confounds"]["names"].get( + "derivatives", {} + ) + if not (isinstance(derivatives_to_compute, dict)): + raise_error( + 'input["confounds"]["names"]["derivatives"] ' + "must be a dictionary. Please check your datagrabber", + ValueError, + ) + + if any( + not (isinstance(k, str) or isinstance(v, str)) + for k, v in derivatives_to_compute.items() + ): + raise_error( + 'input["confounds"]["names"]["derivatives"] ' + "must be a dictionary with string keys and values. " + "Please check your datagrabber", + ValueError, + ) + + missing_derivatives = [ + x + for x in derivatives_to_compute.values() + if x not in confound_df.columns + ] + if len(missing_derivatives) > 0: + raise_error( + "Some of the derivatives to calculate are not in the confounds" + f" dataframe: {missing_derivatives}." + "Please check your data " + f'({input["confounds"]["path"].as_posix()}) ' + "and the datagrabber.", + ValueError, + ) + + t_conf_spec = { + k: input["confounds"]["names"]["spec"][k][v] + for k, v in self.strategy.items() + } + + column_names = set([x for y in t_conf_spec.values() for x in y]) + column_names.add(spike_name) + + missing_columns = [ + x + for x in column_names + if x not in confound_df.columns + and x not in derivatives_to_compute.keys() + ] + + if len(missing_columns) > 0: + raise_error( + "Some of the columns in the confound spec are not in the " + f"confounds dataframe: {missing_columns}. " + "Please check your data " + f'({input["confounds"]["path"].as_posix()}) ' + "and the datagrabber.", + ValueError, + ) + + # TODO: complete type annotations + def fit_transform(self, input): + """Fit and transform.""" + self._validate_data(input) + bold_img = input["BOLD"]["data"] + confounds_df = self._pick_confounds(input["confounds"]) + input["BOLD"]["data"] = self._remove_confounds(bold_img, confounds_df) + + # TODO: Update meta + return input + + +# class FelixConfoundRemover(BaseConfoundRemover): +# """ A class to read confounds from confound files generated by Felix's +# pipeline using CAT and some other custom scripts. It is meant to emulate +# the new nilearn.interfaces.fmriprep.load_confounds as closely as +# possible. + +# """ + +# def read_confounds(self, confound_dataframe): + +# confounds_to_select = [] +# # for some confounds we need to manually calculate derivatives using +# # numpy and add them to the output +# derivatives_to_compute = [] + +# for comp, param in self.strategy.items(): +# if comp == 'motion': +# confounds = [] +# # there should be six rigid body parameters +# for i in range(1, 7): +# # select basic +# confounds_to_select.append(f'RP.{i}') + +# # select squares +# if param in ['power2', 'full']: +# confounds_to_select.append(f'RP^2.{i}') + +# # select derivatives +# if param in ['derivatives', 'full']: +# confounds_to_select.append(f'DRP.{i}') + +# # if 'full' we should not forget the derivative +# # of the squares +# if param in ['full']: +# confounds_to_select.append(f'DRP^2.{i}') + +# elif comp == 'wm_csf': +# confounds = ['WM', 'CSF'] +# elif comp == 'global_signal': +# confounds = ['GS'] + +# for conf in confounds: + +# confounds_to_select.append(conf) + +# # select squares +# if param in ['power2', 'full']: +# confounds_to_select.append(f'{conf}^2') + +# # we have to calculate derivatives (not included in felix' +# # confound files) +# if param in ['derivatives', 'full']: +# derivatives_to_compute.append(conf) + +# if param in ['full']: +# derivatives_to_compute.append(f'{conf}^2') + +# confounds_to_remove = confound_dataframe[confounds_to_select] + +# # calc additional derivatives +# for conf in derivatives_to_compute: +# confounds_to_remove[f'D{conf}'] = np.append( +# np.diff(confound_dataframe[conf]), 0 +# ) + +# # add binary spike regressor if needed at given threshold +# if self.spike is not None: +# fd = confound_dataframe["FD"].copy() +# fd.loc[fd > self.spike] = 1 +# fd.loc[fd != 1] = 0 +# confounds_to_remove['spike'] = fd + +# return confounds_to_remove + + +# class FmriprepConfoundRemover(BaseConfoundRemover): +# """ A ConfoundRemover class for fmriprep output utilising +# nilearn's nilearn.interfaces.fmriprep.load_confounds +# """ + +# def read_confounds(self): +# raise NotImplementedError('read_confounds not implemented') diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py new file mode 100644 index 000000000..8050e1d2d --- /dev/null +++ b/junifer/preprocess/tests/test_confounds.py @@ -0,0 +1,347 @@ +"""Provide tests for confound removal.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import random +import string +from pathlib import Path +from typing import Tuple + +import numpy as np +import pandas as pd +from nibabel import Nifti1Image +from nilearn._utils.niimg_conversions import check_niimg_4d + +from junifer.preprocess.confounds import BaseConfoundRemover + + +# Set RNG seed for reproducibility +np.random.seed(1234567) + + +def generate_conf_name( + size: int = 6, chars: str = string.ascii_uppercase + string.digits +) -> str: + """Generate configuration name.""" + return "".join(random.choice(chars) for _ in range(size)) + + +def _simu_img() -> Tuple[Nifti1Image, Nifti1Image]: + # Random 4D volume with 100 time points + vol = 100 + 10 * np.random.randn(5, 5, 2, 100) + img = Nifti1Image(vol, np.eye(4)) + # Create an nifti image with the data, and corresponding mask + mask = Nifti1Image(np.ones([5, 5, 2]), np.eye(4)) + return img, mask + + +# TODO: split the tests +def test_baseconfoundremover() -> None: + """Test BaseConfoundRemover.""" + # Generate a simulated BOLD img + siimg, simsk = _simu_img() + + # generate random confound dataframe with Felix's column naming + + motion_basic = [f"RP.{i}" for i in range(1, 7)] + motion_power2 = [f"RP^2.{i}" for i in range(1, 7)] + motion_derivatives = [f"DRP.{i}" for i in range(1, 7)] + motion_full = [f"DRP^2.{i}" for i in range(1, 7)] + + wm_csf_basic = ["WM", "CSF"] + wm_csf_power2 = ["WM^2", "CSF^2"] + wm_csf_derivatives = ["DWM", "DCSF"] + wm_csf_full = ["DWM^2", "DCSF^2"] + + gs_basic = ["GS"] + gs_power2 = ["GS^2"] + gs_derivatives = ["DGS"] + gs_full = ["DGS^2"] + + confound_column_names = [] + + confound_column_names.append("FD") # spike + + confound_column_names.extend(motion_basic) + confound_column_names.extend(motion_power2) + confound_column_names.extend(motion_derivatives) + confound_column_names.extend(motion_full) + + confound_column_names.extend(wm_csf_basic) + confound_column_names.extend(wm_csf_power2) + confound_column_names.extend(wm_csf_derivatives) + confound_column_names.extend(wm_csf_full) + + confound_column_names.extend(gs_basic) + confound_column_names.extend(gs_power2) + confound_column_names.extend(gs_derivatives) + confound_column_names.extend(gs_full) + + # add some random irrelevant confounds + for _ in range(10): + confound_column_names.append(generate_conf_name()) + + np.random.shuffle(confound_column_names) + n_cols = len(confound_column_names) + confounds_df = pd.DataFrame( + np.random.randint(0, 100, size=(100, n_cols)), + columns=confound_column_names, + ) + + # Generate spec from Felix's column naming + spec = { + "motion": { + "basic": motion_basic, + "power2": motion_basic + motion_power2, + "derivatives": motion_basic + motion_derivatives, + "full": motion_basic + + motion_derivatives + + motion_power2 + + motion_full, + }, + "wm_csf": { + "basic": wm_csf_basic, + "power2": wm_csf_basic + wm_csf_power2, + "derivatives": wm_csf_basic + wm_csf_derivatives, + "full": wm_csf_basic + + wm_csf_derivatives + + wm_csf_power2 + + wm_csf_full, + }, + "global_signal": { + "basic": gs_basic, + "power2": gs_basic + gs_power2, + "derivatives": gs_basic + gs_derivatives, + "full": gs_basic + gs_derivatives + gs_power2 + gs_full, + }, + } + + # generate a junifer pipeline data object dictionary + input_data_obj = {} + input_data_obj["meta"] = {} + input_data_obj["BOLD"] = {} + input_data_obj["BOLD"]["data"] = siimg + input_data_obj["confounds"] = {} + input_data_obj["confounds"]["path"] = Path("/test.df") + input_data_obj["confounds"]["data"] = confounds_df + input_data_obj["confounds"]["names"] = {} + input_data_obj["confounds"]["names"]["spec"] = spec + input_data_obj["confounds"]["names"]["spike"] = "FD" + + # generate confound removal strategies with varying numbers of parameters + + # Test #1: 36 params, no derivatives to compute, no spike + # 36 params + strat1 = {"motion": "full", "wm_csf": "full", "global_signal": "full"} + + cr = BaseConfoundRemover( + strategy=strat1, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 36 + assert all(x in t_df.columns for x in motion_basic) + assert all(x in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x in t_df.columns for x in wm_csf_power2) + assert all(x in t_df.columns for x in wm_csf_derivatives) + assert all(x in t_df.columns for x in wm_csf_full) + assert all(x in t_df.columns for x in gs_basic) + assert all(x in t_df.columns for x in gs_power2) + assert all(x in t_df.columns for x in gs_derivatives) + assert all(x in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns + + # Test #2: 24 params, no derivatives to compute, no spike + # 24 params + strat2 = { + "motion": "full", + } + + cr = BaseConfoundRemover( + strategy=strat2, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 24 + assert all(x in t_df.columns for x in motion_basic) + assert all(x in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns + + # Test #3: 9 params, no derivatives to compute, no spike + strat3 = {"motion": "basic", "wm_csf": "basic", "global_signal": "basic"} + + cr = BaseConfoundRemover( + strategy=strat3, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 9 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns + + # Test #4: 6 params, no derivatives to compute, no spike + strat4 = { + "motion": "basic", + } + cr = BaseConfoundRemover( + strategy=strat4, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 6 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns + + # Test #5: 2 params, no derivatives to compute, no spike + strat5 = { + "wm_csf": "basic", + } + cr = BaseConfoundRemover( + strategy=strat5, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 2 + assert all(x not in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns + + out = cr.fit_transform(input_data_obj) + check_niimg_4d(out["BOLD"]["data"]) + # TODO: check meta + + # Test #6: 12 params, derivatives to compute, spike + to_select = [ + x for x in confounds_df.columns if x not in motion_derivatives + ] + no_d_df = confounds_df[to_select] + input_data_obj["confounds"]["data"] = no_d_df + + derivatives = {f"D{x}": x for x in motion_basic} + + input_data_obj["confounds"]["names"]["derivatives"] = derivatives + + strat6 = { + "motion": "derivatives", + } + cr = BaseConfoundRemover( + strategy=strat6, spike=0.75, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) + + assert "BOLD" in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj["confounds"]) + assert len(t_df.columns) == 13 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert "FD" not in t_df.columns + assert "spike" in t_df.columns diff --git a/junifer/stats.py b/junifer/stats.py new file mode 100644 index 000000000..f5301d759 --- /dev/null +++ b/junifer/stats.py @@ -0,0 +1,111 @@ +"""Provide functions for statistics.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from functools import partial +from typing import Any, Callable, Dict, Optional + +import numpy as np +from scipy.stats import trim_mean +from scipy.stats.mstats import winsorize + +from .utils import logger, raise_error + + +def get_aggfunc_by_name( + name: str, func_params: Optional[Dict[str, Any]] +) -> Callable: + """Get an aggregation function by its name. + + Parameters + ---------- + name : str + Name to identify the function. Currently supported names and + corresponding functions are: + - 'winsorized_mean' -> scipy.stats.mstats.winsorize + - 'mean' -> numpy.mean + - 'std' -> numpy.std + - 'trim_mean' -> scipy.stats.trim_mean + func_params : dict + Parameters to pass to the function. + E.g. for 'winsorized_mean': func_params = {'limits': [0.1, 0.1]} + + Returns + ------- + function + Respective function with `func_params` parameter set. + + """ + # check validity of names + _valid_func_names = {"winsorized_mean", "mean", "std", "trim_mean"} + if func_params is None: + func_params = {} + # apply functions + if name == "winsorized_mean": + # check validity of func_params + limits = func_params.get("limits") + if limits is None or not isinstance(limits, list): + raise_error( + "func_params must contain a list of limits for " + "winsorized_mean", + ValueError, + ) + if len(limits) != 2: + raise_error( + "func_params must contain a list of two limits for " + "winsorized_mean", + ValueError, + ) + if all((lim >= 0.0 and lim <= 1) for lim in limits): + logger.info(f"Limits for winsorized mean are set to {limits}.") + else: + raise_error( + "Limits for the winsorized mean must be between 0 and 1." + ) + # partially interpret func_params + func = partial(winsorized_mean, **func_params) + elif name == "mean": + func = np.mean + elif name == "std": + func = np.std + elif name == "trim_mean": + if func_params is None: + func = trim_mean + else: + func = partial(trim_mean, **func_params) + else: + raise_error( + f"Function {name} unknown. Please provide any of " + f"{_valid_func_names}" + ) + return func + + +def winsorized_mean( + data: np.ndarray, axis: Optional[int] = None, **win_params +) -> np.ndarray: + """Compute a winsorized mean by chaining winsorization and mean. + + Parameters + ---------- + data : numpy.ndarray + Data to calculate winsorized mean on. + axis : int, optional + The axis to calculate winsorized mean on (default None). + **win_params : dict + Dictionary containing the keyword arguments for the winsorize function. + E.g. {'limits': [0.1, 0.1]} + + Returns + ------- + numpy.ndarray + Winsorized mean of the inputted data with the winsorize settings + applied as specified in win_params. + + """ + win_dat = winsorize(data, axis=axis, **win_params) + win_mean = win_dat.mean(axis=axis) + + return win_mean diff --git a/junifer/storage/__init__.py b/junifer/storage/__init__.py index 4e235bc8b..b8c4f8942 100644 --- a/junifer/storage/__init__.py +++ b/junifer/storage/__init__.py @@ -1,3 +1,9 @@ +"""Provide imports for storage sub-package.""" + # Authors: Federico Raimondo -# Leonard Sasse -# License: AGPL \ No newline at end of file +# Synchon Mandal +# License: AGPL + +from .base import BaseFeatureStorage +from .pandas_base import PandasBaseFeatureStorage +from .sqlite import SQLiteFeatureStorage diff --git a/junifer/storage/base.py b/junifer/storage/base.py new file mode 100644 index 000000000..46a2f0e40 --- /dev/null +++ b/junifer/storage/base.py @@ -0,0 +1,263 @@ +"""Provide abstract base class for feature storage.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from abc import ABC, abstractmethod +from pathlib import Path +from typing import Dict, Iterable, List, Optional, Union + +import pandas as pd + +from .._version import __version__ +from ..utils import raise_error + + +class BaseFeatureStorage(ABC): + """Abstract base class for feature storage. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + Parameters + ---------- + uri : str or pathlib.Path + The path to the storage. + single_output : bool, optional + Whether to have single output (default False). + + """ + + def __init__( + self, uri: Union[str, Path], single_output: bool = False + ) -> None: + """Initialize the class.""" + self.uri = uri + self.single_output = single_output + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + meta : dict + The metadata as a dictionary. + + """ + meta = {} + meta["versions"] = { + "junifer": __version__, + } + return meta + + # TODO: is raising ValueError required? + @abstractmethod + def validate(self, input_: List[str]) -> bool: + """Validate the input to the pipeline step. + + Parameters + ---------- + input_ : list + The input to the pipeline step. + + Returns + ------- + bool + Whether the `input` is valid or not. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + raise_error( + msg="Concrete classes need to implement validate_input().", + klass=NotImplementedError, + ) + + @abstractmethod + def list_features( + self, return_df: bool = False + ) -> Union[Dict[str, Dict], pd.DataFrame]: + """List the features in the storage. + + Parameters + ---------- + return_df : bool, optional + If True, returns a pandas DataFrame. If False, returns a + dictionary (default False). + + Returns + ------- + dict or pandas.DataFrame + List of features in the storage. If dictionary is returned, the + keys are the feature names to be used in read_features() and the + values are the metadata of each feature. + + """ + raise_error( + msg="Concrete classes need to implement list_features().", + klass=NotImplementedError, + ) + + @abstractmethod + def read_df( + self, + feature_name: Optional[str] = None, + feature_md5: Optional[bool] = None, + ) -> pd.DataFrame: + """Read feature from the storage. + + Parameters + ---------- + feature_name : str, optional + Name of the feature to read (default None). + feature_md5 : str, optional + MD5 hash of the feature to read (default None). + + Returns + ------- + pandas.DataFrame + The features as a dataframe. + + """ + raise_error( + msg="Concrete classes need to implement read_df().", + klass=NotImplementedError, + ) + + @abstractmethod + def store_metadata(self, meta: Dict) -> str: + """Store metadata. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + + Returns + ------- + str + The metadata column + + """ + raise_error( + msg="Concrete classes need to implement store_metadata().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_matrix2d( + self, + data, + meta: Dict, + col_names: Optional[Iterable[str]] = None, + row_names: Optional[Iterable[str]] = None, + ) -> None: + """Store 2D matrix. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + col_names : list or tuple of str, optional + The column names (default None). + row_names : list of tuple of str, optional + The row names (default None). + + """ + raise_error( + msg="Concrete classes need to implement store_matrix2d().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_table( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: Optional[str] = None, + ) -> None: + """Store table. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name to use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + raise_error( + msg="Concrete classes need to implement store_table().", + klass=NotImplementedError, + ) + + @abstractmethod + def store_df(self, df: pd.DataFrame, meta: Dict) -> None: + """Store pandas DataFerame. + + Parameters + ---------- + df : pandas.DataFrame + The DataFrame to store. + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="Concrete classes need to implement store_df().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_timeseries(self, data, meta: Dict) -> None: + """Store timeseries. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="Concrete classes need to implement store_timeseries().", + klass=NotImplementedError, + ) + + @abstractmethod + def collect(self) -> None: + """Collect data.""" + raise_error( + msg="Concrete classes need to implement collect().", + klass=NotImplementedError, + ) + + def __str__(self) -> str: + """Represent object as string. + + Returns + ------- + str + The string representation. + + """ + single = ( + "(single output)" + if self.single_output is True + else "(multiple output)" + ) + return f"<{self.__class__.__name__} @ {self.uri} {single}>" diff --git a/junifer/storage/pandas_base.py b/junifer/storage/pandas_base.py new file mode 100644 index 000000000..eaa2a0a32 --- /dev/null +++ b/junifer/storage/pandas_base.py @@ -0,0 +1,65 @@ +"""Provide abstract base class for feature storage via pandas.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import json +from pathlib import Path +from typing import Dict, Union + +import pandas as pd + +from .base import BaseFeatureStorage + + +class PandasBaseFeatureStorage(BaseFeatureStorage): + """Abstract base class for feature storage via pandas. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + Parameters + ---------- + uri : str or pathlib.Path + The path to the storage. + single_output : bool, optional + Whether to have single output (default False). + **kwargs + Keyword arguments passed to superclass. + + See Also + -------- + BaseFeatureStorage + + """ + + def __init__( + self, uri: Union[str, Path], single_output: bool = False, **kwargs + ) -> None: + """Initialize the class.""" + super().__init__(uri=uri, single_output=single_output, **kwargs) + + def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame: + """Convert the metadata to a pandas DataFrame. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + meta_md5 : str + The MD5 hash of the metadata. + + Returns + ------- + pandas.DataFrame + + """ + data_df = {} + for k, v in meta.items(): + data_df[k] = json.dumps(v, sort_keys=True) + if "marker" in meta: + data_df["name"] = meta["marker"]["name"] + df = pd.DataFrame(data_df, index=[meta_md5]) + df.index.name = "meta_md5" + return df diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py new file mode 100644 index 000000000..9522713b0 --- /dev/null +++ b/junifer/storage/sqlite.py @@ -0,0 +1,598 @@ +"""Provide concrete implementation for feature storage via SQLite.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import TYPE_CHECKING, Dict, Iterable, List, Optional, Union + +import pandas as pd +from pandas.core.base import NoNewAttributesMixin +from pandas.io.sql import pandasSQL_builder +from sqlalchemy import create_engine, inspect +from tqdm import tqdm + +from ..api.decorators import register_storage +from ..utils import logger, raise_error, warn_with_log +from .pandas_base import PandasBaseFeatureStorage +from .utils import element_to_index, element_to_prefix, process_meta + + +if TYPE_CHECKING: + from sqlalchemy.engine import Engine + + +@register_storage +class SQLiteFeatureStorage(PandasBaseFeatureStorage): + """Concrete implementation for feature storage via SQLite. + + Parameters + ---------- + uri : str or pathlib.Path + The path to the file to be used. + single_output : bool, optional + If False, will create one file per element. The name + of the file will be prefixed with the respective element. + If True, will create only one file as specified in the `uri` and + store all the elements in the same file. This behaviour is only + suitable for non-parallel executions. SQLite does not support + concurrency (default False). + upsert : {"ignore", "update"}, optional + Upsert mode. If "ignore" is used, the existing elements are ignored. + If "update", the existing elements are updated (default "update"). + **kwargs : dict + The keyword arguments passed to the superclass. + + See Also + -------- + PandasBaseFeatureStorage + + """ + + def __init__( + self, + uri: Union[str, Path], + single_output: bool = False, + upsert: str = "update", + **kwargs: str, + ) -> None: + """Initialize the class.""" + if upsert not in ["update", "ignore"]: + raise_error( + msg=( + "Invalid choice for `upsert`. " + "Must be either 'update' or 'ignore'." + ) + ) + # Convert str to Path + if not isinstance(uri, Path): + uri = Path(uri) + # Create parent directories if not present + if not uri.parent.exists(): + logger.info( + f"Output directory ({str(uri.parent.absolute())}) " + "does not exist, creating now." + ) + uri.parent.mkdir(parents=True, exist_ok=True) + super().__init__(uri=uri, single_output=single_output, **kwargs) + self._upsert = upsert + self._valid_inputs = ["table", "timeseries"] + + def get_engine(self, meta: Optional[Dict] = None) -> "Engine": + """Get engine. + + Parameters + ---------- + meta : dict, optional + The metadata as dictionary (default None). + + Returns + ------- + sqlalchemy.Engine + The sqlalchemy engine. + + """ + # Set metadata as empty dictionary if None + if meta is None: + meta = {} + # Retrieve element key from metadata + element = meta.get("element", None) + # Prefixed elements + prefix = "" + if self.single_output is False: + if element is None: + raise_error( + msg="element must be specified when" + "single_output is False." + ) + prefix = element_to_prefix(element) + # Format URI for engine creation + uri = ( + "sqlite:///" f"{self.uri.parent}/{prefix}{self.uri.name}" + ) # type: ignore + return create_engine(uri, echo=False) + + def _save_upsert( + self, + df: pd.DataFrame, + name: str, + engine: Optional["Engine"] = None, + if_exists: str = "append", + ) -> None: + """Implement UPSERT functionality. + + Parameters + ---------- + df : pandas.DataFrame + DataFrame to save. + name : str + Name of the table to save. + engine : sqlalchemy.Engine, optional + The sqlalchemy engine to use (default None). + if_exists : {"replace", "nocheck", "append", "fail"}, optional + Action to take if the table exists. If "replace", existing table + will be dropped before inserting new values. If "nocheck", + existing table will be ignored. If "append", the data will be + appended to the existing table. If "fail", it will raise an error + (default "append"). + + Raises + ------ + ValueError + If the table exists and if_exists is "fail" or if invalid option is + passed to `if_exists`. + + """ + # Get index names + index_col = df.index.names + # Get sqlalchemy engine if None + if engine is None: + engine = self.get_engine() + # Write data + with engine.begin() as con: + # Check for table's existence + if not inspect(engine).has_table(name): + # New table, so no big issue + df.to_sql(name=name, con=con, if_exists="append") + else: + if if_exists == "replace": + # Replace all the existing elements + df.to_sql(name=name, con=con, if_exists="replace") + elif if_exists == "nocheck": + # Ignore check + df.to_sql(name, con=con, if_exists="append") + elif if_exists == "append": + # TODO: improve + # Step 1: split incoming data into existing and new data + pk_indb = _get_existing_pk( + con, table_name=name, index_col=index_col + ) + existing, new = _split_incoming_data( + df, pk_indb, index_col + ) + # Step 2: upsert existing data + pandas_sql = pandasSQL_builder(con) + pandas_sql.meta.reflect(bind=con, only=[name]) + table = pandas_sql.get_table(name) + update_stmts = NoNewAttributesMixin + if len(existing) > 0 and len(new) > 0: + warn_with_log( + f"Some rows (n={len(existing)}) are already " + "present in the database. The storage is " + f"configured to {self._upsert} the existing " + f"elements. The new rows (n={len(new)}) will be " + "appended. This warning is shown because normally " + "all of the elements should be updated." + ) + if self._upsert == "update": + update_stmts = _generate_update_statements( + table, index_col, existing + ) + for stmt in update_stmts: + con.execute(stmt) + # Step 3: insert new data + new.to_sql(name=name, con=con, if_exists="append") + elif if_exists == "fail": + # Case 4: existing table, so we need to check if the index + # is present or not. + raise_error(msg=f"Table ({name}) already exists.") + else: + raise_error( + msg=f"Invalid option {if_exists} for if_exists." + ) + + # TODO: complete type annotations + def store_2d( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: Optional[str] = None, + ) -> None: + """Store 2D dataframe. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name to use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + n_rows = len(data) + # Convert element metadata to index + idx = element_to_index( + meta=meta, n_rows=n_rows, rows_col_name=rows_col_name + ) + # Prepare new dataframe + data_df = pd.DataFrame( + data, columns=columns, index=idx + ) # type: ignore + # Store dataframe + self.store_df(df=data_df, meta=meta) + + def validate(self, input_: List[str]) -> bool: + """Implement input validation. + + Parameters + ---------- + input_ : list of str + The input to the pipeline step. + + Returns + ------- + bool + Whether the `input` is valid or not. + + """ + # Convert input to list + if not isinstance(input_, list): + input_ = [input_] + + return all(x in self._valid_inputs for x in input_) + + def list_features( + self, return_df: bool = False + ) -> Union[Dict[str, Dict], pd.DataFrame]: + """Implement features listing from the storage. + + Parameters + ---------- + return_df : bool, optional + If True, returns a pandas DataFrame. If False, returns a + dictionary (default False). + + Returns + ------- + dict or pandas.DataFrame + List of features in the storage. If dictionary is returned, the + keys are the feature names to be used in read_features() and the + values are the metadata of each feature. + + """ + meta_df = pd.read_sql( + sql="meta", + con=self.get_engine(), + index_col="meta_md5", + ) + out = meta_df + # Return dictionary + if return_df is False: + out = meta_df.to_dict(orient="index") + return out + + def read_df( + self, + feature_name: Optional[str] = None, + feature_md5: Optional[str] = None, + ) -> pd.DataFrame: + """Implement feature reading from the storage. + + Either one of `feature_name` or `feature_md5` needs to be specified. + + Parameters + ---------- + feature_name : str, optional + Name of the feature to read (default None). + feature_md5 : str, optional + MD5 hash of the feature to read (default None). + + Returns + ------- + pandas.DataFrame + The features as a dataframe. + + Raises + ------ + ValueError + If parameter values are invalid or feature is not found or + multiple features are found. + + """ + # Get sqlalchemy engine + engine = self.get_engine() + # Parameter value check + if feature_md5 is not None and feature_name is not None: + raise_error( + msg=( + "Only one of `feature_name` or `feature_md5` can be " + "specified." + ) + ) + elif feature_md5 is None and feature_name is None: + raise_error( + msg=( + "At least one of `feature_name` or `feature_md5` " + "must be specified." + ) + ) + elif feature_md5 is not None: + table_name = f"meta_{feature_md5}" + else: + meta_df = pd.read_sql( + sql="meta", + con=engine, + index_col="meta_md5", + ) + t_df = meta_df.query(f"name == '{feature_name}'") + if len(t_df) == 0: + raise_error(msg=f"Feature {feature_name} not found") + elif len(t_df) > 1: + raise_error( + msg=( + f"More than one feature with name {feature_name} " + "found. This file is invalid. You can bypass this " + "issue by specifying a `feature_md5`." + ) + ) + table_name = f"meta_{t_df.index[0]}" + # Read metadata from table + df = pd.read_sql(sql=table_name, con=engine) + # Read the index + query = ( + "SELECT ii.name FROM sqlite_master AS m, " + "pragma_index_list(m.name) AS il, " + "pragma_index_info(il.name) AS ii " + f"WHERE tbl_name='{table_name}' " + "ORDER BY cid;" + ) + index_names = ( + pd.read_sql(sql=query, con=engine).values.squeeze().tolist() + ) + # Set index on dataframe + df = df.set_index(index_names) + return df + + def store_metadata(self, meta: Dict) -> str: + r"""Implement metadata storing in the storage. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + + Returns + ------- + str + The MD5 hash of the metadata prefixed with "meta\_". + + """ + # Copy metadata + t_meta = meta.copy() + # Update metadata + t_meta.update(self.get_meta()) + # Process metadata + meta_md5, t_meta_row = process_meta(t_meta) + # Get sqlalchemy engine + engine = self.get_engine(meta=t_meta) + if meta_md5 not in inspect(engine).get_table_names(): + # Convert metadata to dataframe + meta_df = self._meta_row(meta=t_meta_row, meta_md5=meta_md5) + # Save dataframe + self._save_upsert(meta_df, "meta", engine) + return f"meta_{meta_md5}" + + # TODO: complete type annotations + def store_matrix2d( + self, + data, + meta: Dict, + col_names: Optional[Iterable[str]] = None, + rows_col_name: Optional[str] = None, + ) -> None: + """Implement 2D matrix storing. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + col_names : list or tuple of str, optional + The column names (default None). + rows_col_name : str, optional + The column name to use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + # Same as store_2d, but order is important + raise_error( + msg="store_matrix2d() not implemented", klass=NotImplementedError + ) + + # TODO: complete type annotations + def store_table( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: Optional[str] = None, + ) -> None: + """Implement table storing. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name to use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + self.store_2d( + data=data, meta=meta, columns=columns, rows_col_name=rows_col_name + ) + + def store_df(self, df: pd.DataFrame, meta: Dict) -> None: + """Implement dataframe storing. + + Parameters + ---------- + df : pandas.DataFrame + The DataFrame to store. + meta : dict + The metadata as a dictionary. + + Raises + ------ + ValueError + If the dataframe index has items that are not in the index + generated from the metadata. + + """ + # TODO: Test this function + # Check that the index generated by meta matches the one in + # the dataframe. + idx = element_to_index(meta) + # Given the meta, we might not know if there is an extra column added + # when storing a timeseries or 2d elements. We need to check if the + # extra element is only one. + extra = [x for x in df.index.names if x not in idx.names] + if len(extra) > 1: + raise_error( + "The index of the dataframe has extra items that are not " + "in the index generated from the metadata." + ) + elif len(extra) == 1: + # The df has one extra index item, this should be the new name + # of the missing element in the index + idx = element_to_index(meta, rows_col_name=extra[0]) + + if any(x not in df.index.names for x in idx.names): + raise_error( + "The index of the dataframe is missing index items that are " + "generated from the metadata." + ) + # Get table name + table_name = self.store_metadata(meta) + # Get sqlalchemy engine + engine = self.get_engine(meta) + # Save data + self._save_upsert(df, table_name, engine) + + # TODO: complete type annotations + def store_timeseries(self, data, meta: Dict) -> None: + """Implement timeseries storing. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="store_timeseries() not implemented.", + klass=NotImplementedError, + ) + + def collect(self) -> None: + """Implement data collection. + + Raises + ------ + NotImplementedError + If `single_output` is True. + + """ + if self.single_output is True: + raise_error(msg="collect() is not implemented for single output.") + logger.info( + "Collecting data from " f"{self.uri.parent}/*{self.uri.name}" + ) # type: ignore + # Create new instance + out_storage = SQLiteFeatureStorage( + uri=self.uri, single_output=True, upsert="ignore" + ) + # Glob files + files = self.uri.parent.glob(f"*{self.uri.name}") # type: ignore + for elem in tqdm(files, desc="file"): + logger.debug(f"Reading from {str(elem.absolute())}") + in_storage = SQLiteFeatureStorage(uri=elem, single_output=True) + in_engine = in_storage.get_engine() + # Open "meta" table + t_meta_df = pd.read_sql( + sql="meta", con=in_engine, index_col="meta_md5" + ) + # Save metadata + out_storage._save_upsert(t_meta_df, "meta") + # Save dataframes + for meta_md5 in tqdm(t_meta_df.index, desc="feature"): + logger.debug(f"Collecting feature {meta_md5}") + # TODO: Fix this, needs that read_feature sets the index + # properly + table_name = f"meta_{meta_md5}" + t_df = in_storage.read_df(feature_md5=meta_md5) + # Save data + out_storage._save_upsert(t_df, table_name, if_exists="nocheck") + + +# TODO: refactor +def _get_existing_pk(con, table_name, index_col): + pk_cols = ", ".join(index_col) + query = f"SELECT {pk_cols} FROM {table_name};" + pk_indb = pd.read_sql(query, con=con) + return pk_indb + + +# TODO: refactor +def _split_incoming_data(df, pk_indb, index_col): + incoming_pk = df.reset_index()[index_col] + exists_mask = ( + incoming_pk[index_col] + .apply(tuple, axis=1) + .isin(pk_indb[index_col].apply(tuple, axis=1)) + ) + existing, new = df.loc[exists_mask.values], df.loc[~exists_mask.values] + return existing, new + + +# TODO: refactor +def _generate_update_statements(table, index_col, rows_to_update): + from sqlalchemy import and_ + + new_records = rows_to_update.to_dict(orient="records") + pk_indb = rows_to_update.reset_index()[index_col] + pk_cols = [table.c[key] for key in index_col] + + stmts = [] + for i, (_, keys) in enumerate(pk_indb.iterrows()): + stmt = ( + table.update() + .where( + and_(col == keys[j] for j, col in enumerate(pk_cols)) + ) # type: ignore + .values(new_records[i]) + ) + stmts.append(stmt) + return stmts diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py new file mode 100644 index 000000000..af5c4a3e0 --- /dev/null +++ b/junifer/storage/tests/test_sqlite.py @@ -0,0 +1,589 @@ +"""Provide tests for sqlite.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import List, Union + +import numpy as np +import pandas as pd +import pytest +from pandas.testing import assert_frame_equal +from sqlalchemy import create_engine + +from junifer.storage.sqlite import SQLiteFeatureStorage +from junifer.storage.utils import ( + element_to_index, + element_to_prefix, + process_meta, +) + + +df1 = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5], + "pk2": ["a", "b", "c", "d", "e"], + "col1": [11, 22, 33, 44, 55], + "col2": [111, 222, 333, 444, 555], + } +).set_index(["element", "pk2"]) + +df2 = pd.DataFrame( + { + "element": [2, 5, 6], + "pk2": ["b", "e", "f"], + "col1": [2222, 5555, 66], + "col2": [22222, 55555, 666], + } +).set_index(["element", "pk2"]) + +df_update = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5, 6], + "pk2": ["a", "b", "c", "d", "e", "f"], + "col1": [11, 2222, 33, 44, 5555, 66], + "col2": [111, 22222, 333, 444, 55555, 666], + } +).set_index(["element", "pk2"]) + +df_ignore = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5, 6], + "pk2": ["a", "b", "c", "d", "e", "f"], + "col1": [11, 22, 33, 44, 55, 66], + "col2": [111, 222, 333, 444, 555, 666], + } +).set_index(["element", "pk2"]) + + +def _read_sql( + table_name: str, uri: str, index_col: Union[str, List[str]] +) -> pd.DataFrame: + """Read database table into a pandas DataFrame. + + Parameters + ---------- + table_name : str + The table name. + uri : str + The URI of the database. + index_col : str + The index column name. + + Returns + ------- + pandas.DataFrame + The contents of the table in a DataFrame. + + """ + engine = create_engine(f"sqlite:///{uri}", echo=False) + df = pd.read_sql(sql=table_name, con=engine, index_col=index_col) + return df + + +def test_get_engine_single_output(tmp_path: Path) -> None: + """Test engine retrieval with single output. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_single_output.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + assert storage.single_output is True + engine = storage.get_engine() + assert engine.url.drivername == "sqlite" + assert engine.url.database == str(uri.absolute()) + + +def test_get_engine_multi_output(tmp_path: Path) -> None: + """Test engine retrieval with multi output. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_multi_output.db" + storage = SQLiteFeatureStorage( + uri=uri, single_output=False, upsert="ignore" + ) + with pytest.raises(ValueError, match="element must be specified"): + storage.get_engine() + + +def test_get_engine_single_output_creation(tmp_path: Path) -> None: + """Test engine retrieval with single output creation. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + tocreate = tmp_path / "tocreate" + # Path does not exist yet + assert not tocreate.exists() + uri = tocreate.absolute() / "test_single_output.db" + _ = SQLiteFeatureStorage(uri=uri, single_output=True, upsert="ignore") + # Path exists now + assert tocreate.exists() + + +def test_upsert_replace(tmp_path: Path) -> None: + """Test dataframe store with upsert=replace. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_upsert_replace.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df=df1, meta=meta) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Upsert using replace + storage._save_upsert(df=df2, name=table_name, if_exists="replace") + # Read stored table + c_df2 = _read_sql( + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df2, c_df2) + + +def test_upsert_ignore(tmp_path: Path) -> None: + """Test dataframe store with upsert=ignore. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_upsert_ignore.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df=df1, meta=meta) + # Store metadata + table_name = storage.store_metadata(meta=meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Check for warning + with pytest.warns(RuntimeWarning, match="are already present"): + storage.store_df(df2, meta) + # Read stored table + c_dfignore = _read_sql( + table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(c_dfignore, df_ignore) + # Check for error + with pytest.raises(ValueError, match=r"already exists"): + storage._save_upsert(df2, table_name, if_exists="fail") + + +def test_upsert_update(tmp_path: Path) -> None: + """Test dataframe store with upsert=delete. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_upsert_delete.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df1, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Save to database + storage.store_df(df2, meta) + # Read stored table + c_dfupdate = _read_sql( + table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(c_dfupdate, df_update) + + +def test_upsert_invalid_option(tmp_path: Path) -> None: + """Test dataframe store with invalid option for upsert. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_upsert_invalid.db" + with pytest.raises(ValueError): + SQLiteFeatureStorage(uri=uri, single_output=True, upsert="wrong") + + +# TODO: can the tests be separated? +def test_store_df_and_read_df(tmp_path: Path) -> None: + """Test storing dataframe and reading of stored table into dataframe. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_df_and_read_df.db" + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = { + "element": "test", + "version": "0.0.1", + "marker": {"name": "fcname"}, + } + # Columns to store + to_store = df1[["col1", "col2"]] + # Check for error while storing + with pytest.raises(ValueError, match=r"missing index items"): + storage.store_df(to_store.set_index("col1"), meta) + # Set index + to_store = df1.reset_index().set_index(["element", "pk2", "col1"]) + # Check for error while storing + with pytest.raises(ValueError, match=r"extra items"): + storage.store_df(to_store, meta) + # Convert element to index + idx = element_to_index(meta, n_rows=len(to_store)) + # Set index + to_store = to_store.set_index(idx) + # Store dataframe + storage.store_df(to_store, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # List stored features + features = storage.list_features() + # Check correct usage + assert len(features) == 1 + assert table_name.replace("meta_", "") in features + # Check for missing feature + with pytest.raises(ValueError, match="not found"): + storage.read_df("wrong_md5") + # Check for missing feature to fetch + with pytest.raises(ValueError, match="least one"): + storage.read_df() + # Check for multiple features to fetch + with pytest.raises(ValueError, match="Only one"): + storage.read_df("wrong_md5", "wrong_name") + # Get MD5 hash of features + feature_md5 = list(features.keys())[0] + # Check for key + assert "fcname" == features[feature_md5]["name"] + # Read into dataframes + read_df1 = storage.read_df(feature_md5=feature_md5) + read_df2 = storage.read_df(feature_name="fcname") + # Check if dataframes are equal + assert_frame_equal(read_df1, read_df2) + assert_frame_equal(read_df1, to_store) + + +def test_store_metadata(tmp_path: Path) -> None: + """Test metadata store. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_metadata_store.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Store metadata + table_name = storage.store_metadata(meta) + assert table_name.startswith("meta_") + + +def test_store_table(tmp_path: Path) -> None: + """Test table store. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_table.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + # Metadata to store + meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} + # Data to store + data = [ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ] + # Convert element to index + idx = element_to_index(meta, n_rows=5, rows_col_name="scan") + # Create dataframe + df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) + # Store table + storage.store_table(data, meta, columns=["f1", "f2"], rows_col_name="scan") + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df = _read_sql( + table_name=table_name, + uri=uri.as_posix(), + index_col=["element", "scan"], + ) + # Check if dataframes are equal + assert_frame_equal(df, c_df) + + # New data to store + data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] + # Convert element to index + idx_new = element_to_index(meta, n_rows=6, rows_col_name="scan") + # Create dataframe + df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new) + # Check warning + with pytest.warns(RuntimeWarning, match=r"Some rows"): + storage.store_table( + data_new, meta, columns=["f1", "f2"], rows_col_name="scan" + ) + # Read stored table + c_df_new = _read_sql( + table_name=table_name, + uri=uri.as_posix(), + index_col=["element", "scan"], + ) + # Check if dataframes are equal + assert_frame_equal(df_new, c_df_new) + + +# TODO: can the test be parametrized? +def test_store_multiple_output(tmp_path: Path): + """Test storing using single_output=False. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_multiple_output.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=False) + # Metadata to store + meta1 = { + "element": {"subject": "test-01", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta2 = { + "element": {"subject": "test-02", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta3 = { + "element": {"subject": "test-01", "session": "ses-02"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + # Data to store + data1 = np.array( + [ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ] + ) + data2 = data1 * 10 + data3 = data1 * 20 + # Process metadata for storage + hash1, _ = process_meta(meta1) + # Convert element to index + idx1 = element_to_index(meta1, n_rows=5, rows_col_name="scan") + # Create dataframe + df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1) + # Process metadata for storage + hash2, _ = process_meta(meta2) + # Convert element to index + idx2 = element_to_index(meta2, n_rows=5, rows_col_name="scan") + # Create dataframe + df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2) + # Process metadata for storage + hash3, _ = process_meta(meta3) + # Convert element to index + idx3 = element_to_index(meta3, n_rows=5, rows_col_name="scan") + # Create dataframe + df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3) + # Check hash equality + assert hash1 == hash2 + assert hash2 == hash3 + # Store tables + storage.store_table( + data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data2, meta2, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data3, meta3, columns=["f1", "f2"], rows_col_name="scan" + ) + # Check that URI does not exist yet + assert not uri.exists() + # Convert element to preifx + prefix1 = element_to_prefix(meta1["element"]) + prefix2 = element_to_prefix(meta2["element"]) + prefix3 = element_to_prefix(meta3["element"]) + # URIs for data storage + uri1 = uri.parent / f"{prefix1}{uri.name}" + uri2 = uri.parent / f"{prefix2}{uri.name}" + uri3 = uri.parent / f"{prefix3}{uri.name}" + # Check URIs for data storage exist + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + # Store metadata + table_name = storage.store_metadata(meta1) + # Set index columns + cols = ["subject", "session", "scan"] + # Read stored tables + cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) + cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols) + cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=cols) + # Check if dataframes are equal + assert_frame_equal(df1, cdf1) + assert_frame_equal(df2, cdf2) + assert_frame_equal(df3, cdf3) + + +# TODO: can test be paramtrized? +def test_collect(tmp_path: Path) -> None: + """Test collect. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_collect.db" + storage = SQLiteFeatureStorage(uri=uri) + # Metadata for storage + meta1 = { + "element": {"subject": "test-01", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta2 = { + "element": {"subject": "test-02", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta3 = { + "element": {"subject": "test-01", "session": "ses-02"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + # Data for storage + data1 = np.array( + [ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ] + ) + data2 = data1 * 10 + data3 = data1 * 20 + # Store tables + storage.store_table( + data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data2, meta2, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data3, meta3, columns=["f1", "f2"], rows_col_name="scan" + ) + # Convert element to prefix + prefix1 = element_to_prefix(meta1["element"]) + prefix2 = element_to_prefix(meta2["element"]) + prefix3 = element_to_prefix(meta3["element"]) + # URIs for data storage + uri1 = uri.parent / f"{prefix1}{uri.name}" + uri2 = uri.parent / f"{prefix2}{uri.name}" + uri3 = uri.parent / f"{prefix3}{uri.name}" + # Check URIs for data storage exist + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + # Check that URI does not exist yet + assert not uri.exists() + # Collect data + storage.collect() + # Check that URI exists now + assert uri.exists() + # Set index columns + cols = ["subject", "session", "scan"] + # Store metadata + table_name = storage.store_metadata(meta1) + # Read stored tables + all_df = _read_sql(table_name, uri.as_posix(), index_col=cols) + cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) + cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols) + cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=cols) + # Operate on retrieved tables + all_cdf = pd.concat([cdf1, cdf2, cdf3]) + all_df.sort_index(level=cols, inplace=True) + all_cdf.sort_index(level=cols, inplace=True) + # Check if dataframes are equal + assert_frame_equal(all_df, all_cdf) diff --git a/junifer/storage/tests/test_storage_base.py b/junifer/storage/tests/test_storage_base.py new file mode 100644 index 000000000..c9af9db04 --- /dev/null +++ b/junifer/storage/tests/test_storage_base.py @@ -0,0 +1,85 @@ +"""Provide tests for base.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import pytest + +from junifer.storage.base import BaseFeatureStorage + + +def test_BaseFeatureStorage_abstractness() -> None: + """Test BaseFeatureStorage is abstract base class.""" + with pytest.raises(TypeError, match=r"abstract"): + BaseFeatureStorage(uri="/tmp") # type: ignore + + +def test_BaseFeatureStorage() -> None: + """Test BaseFeatureStorage.""" + # Create concrete class + class MyFeatureStorage(BaseFeatureStorage): + def __init__(self, uri, single_output=False): + super().__init__(uri, single_output=single_output) + + def validate(self, input): + super().validate(input) + + def list_features(self): + super().list_features() + + def read_df(self, feature_name=None, feature_md5=None): + super().read_df(feature_name=feature_name, feature_md5=feature_md5) + + def store_metadata(self, metadata): + super().store_metadata(metadata) + + def store_matrix2d(self, matrix, meta): + super().store_matrix2d(matrix, meta) + + def store_table(self, table, meta): + super().store_table(table, meta) + + def store_df(self, df, meta): + super().store_df(df, meta) + + def store_timeseries(self, timeseries, meta): + super().store_timeseries(timeseries, meta) + + def collect(self): + return super().collect() + + st = MyFeatureStorage(uri="/tmp") + assert st.single_output is False + + st = MyFeatureStorage(uri="/tmp", single_output=True) + assert st.single_output is True + + with pytest.raises(NotImplementedError): + st.validate(None) + + with pytest.raises(NotImplementedError): + st.list_features() + + with pytest.raises(NotImplementedError): + st.read_df(None) + + with pytest.raises(NotImplementedError): + st.store_metadata(None) + + with pytest.raises(NotImplementedError): + st.store_matrix2d(None, None) + + with pytest.raises(NotImplementedError): + st.store_table(None, None) + + with pytest.raises(NotImplementedError): + st.store_df(None, None) # type: ignore + + with pytest.raises(NotImplementedError): + st.store_timeseries(None, None) + + with pytest.raises(NotImplementedError): + st.collect() + + assert st.uri == "/tmp" diff --git a/junifer/storage/tests/test_utils.py b/junifer/storage/tests/test_utils.py new file mode 100644 index 000000000..a5b2b466d --- /dev/null +++ b/junifer/storage/tests/test_utils.py @@ -0,0 +1,200 @@ +"""Provide tests for utils.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List, Tuple, Union + +import pytest + +from junifer.storage.utils import ( + element_to_index, + element_to_prefix, + process_meta, +) + + +def test_process_meta_invalid_metadata_type() -> None: + """Test invalid metadata type check for metadata hash processing.""" + meta = None + with pytest.raises(ValueError, match=r"`meta` must be a dict"): + process_meta(meta) # type: ignore + + +# TODO: parameterize +def test_process_meta_hash() -> None: + """Test metadata hash processing.""" + meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} + hash1, _ = process_meta(meta) + + meta = {"element": "foo", "B": [2, 3, 4, 5, 6], "A": 1} + hash2, _ = process_meta(meta) + assert hash1 == hash2 + + meta = {"element": "foo", "A": 1, "B": [2, 3, 1, 5, 6]} + hash3, _ = process_meta(meta) + assert hash1 != hash3 + + meta1 = { + "element": "foo", + "B": { + "B2": [2, 3, 4, 5, 6], + "B1": [9.22, 3.14, 1.41, 5.67, 6.28], + "B3": (1, "car"), + }, + "A": 1, + } + + meta2 = { + "A": 1, + "B": { + "B3": (1, "car"), + "B1": [9.22, 3.14, 1.41, 5.67, 6.28], + "B2": [2, 3, 4, 5, 6], + }, + "element": "foo", + } + + hash4, _ = process_meta(meta1) + hash5, _ = process_meta(meta2) + assert hash4 == hash5 + + +def test_process_meta_invalid_metadata_key() -> None: + """Test invalid metadata key check for metadata hash processing.""" + meta = {} + with pytest.raises(ValueError, match=r"_element_keys"): + process_meta(meta) + + +@pytest.mark.parametrize( + "meta,elements", + [ + ({"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}, ["element"]), + ( + { + "element": {"subject": "foo", "session": "bar"}, + "B": [2, 3, 4, 5, 6], + "A": 1, + }, + ["subject", "session"], + ), + ], +) +def test_process_meta_element(meta: Dict, elements: List[str]) -> None: + """Test metadata element after processing. + + Parameters + ---------- + meta : dict + The parametrized metadata dictionary. + elements : list of str + The parametrized elements to assert against. + + """ + _, processed_meta = process_meta(meta) + assert "_element_keys" in processed_meta + assert processed_meta["_element_keys"] == elements + assert "A" in processed_meta + assert "B" in processed_meta + + +@pytest.mark.parametrize( + "element,prefix", + [ + ("sub-01", "element_sub-01_"), + (1, "element_1_"), + ({"subject": "sub-01"}, "element_sub-01_"), + ({"subject": 1}, "element_1_"), + ({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"), + ({"subject": 1, "session": 2}, "element_1_2_"), + (("sub-01", "ses-02"), "element_sub-01_ses-02_"), + ((1, 2), "element_1_2_"), + ], +) +def test_element_to_prefix( + element: Union[str, int, Dict, Tuple], prefix: str +) -> None: + """Test converting element to prefix (for file naming). + + Parameters + ---------- + element : str, int, dict or tuple + The parameterized element. + prefix : str + The parametrized prefix to assert against. + + """ + prefix_generated = element_to_prefix(element) + assert prefix_generated == prefix + + +def test_element_to_index_check_meta_invalid_key() -> None: + """Test element to index metadata key checking.""" + meta = {"noelement": "foo"} + with pytest.raises(ValueError, match=r"metadata must contain the key"): + element_to_index(meta) + + +def test_element_to_index() -> None: + """Test element to index.""" + meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} + index = element_to_index(meta) + assert index.names == ["element", "idx"] + assert index.levels[0].name == "element" + assert index.levels[0].values[0] == "foo" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) + + index = element_to_index(meta, n_rows=10) + assert index.names == ["element", "idx"] + assert index.levels[0].name == "element" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (10,) + + index = element_to_index(meta, n_rows=1, rows_col_name="scan") + assert index.names == ["element", "scan"] + assert index.levels[0].name == "element" + assert index.levels[0].values[0] == "foo" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == "scan" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) + + index = element_to_index(meta, n_rows=7, rows_col_name="scan") + assert index.names == ["element", "scan"] + assert index.levels[0].name == "element" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "scan" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (7,) + + meta = { + "element": {"subject": "sub-01", "session": "ses-01"}, + "A": 1, + "B": [2, 3, 4, 5, 6], + } + index = element_to_index(meta, n_rows=10) + + assert index.levels[0].name == "subject" + assert all(x == "sub-01" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "session" + assert all(x == "ses-01" for x in index.levels[1].values) + assert index.levels[1].values.shape == (1,) + + assert index.levels[2].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[2].values)) + assert index.levels[2].values.shape == (10,) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py new file mode 100644 index 000000000..0a6fa58df --- /dev/null +++ b/junifer/storage/utils.py @@ -0,0 +1,164 @@ +"""Provide utility functions for the storage sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import hashlib +import json +from typing import Any, Dict, Optional, Tuple, Union + +import numpy as np +import pandas as pd + +from ..utils.logging import logger, raise_error + + +def _meta_hash(meta: Dict) -> str: + """Compute the MD5 hash of the metadata. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element". + + Returns + ------- + str + The MD5 hash of the metadata. + + """ + logger.debug(f"Hashing metadata: {meta}") + meta_md5 = hashlib.md5( + json.dumps(meta, sort_keys=True).encode("utf-8") + ).hexdigest() + logger.debug(f"Hash computed: {meta_md5}") + return meta_md5 + + +def process_meta(meta: Dict) -> Tuple[str, Dict]: + """Process the metadata for storage. + + It removes the key "element" and adds the "_element_keys" with the keys + used to index the element. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element". + + Returns + ------- + str + The MD5 hash of the metadata. + dict + The processed metadata for storage. + + Raises + ------ + ValueError + If `meta` is None or if it does not contain the key "element" or + "_element_keys". + + """ + if meta is None: + raise_error(msg="`meta` must be a dict (currently is None)") + # Copy the metadata + t_meta = meta.copy() + # Remove key "element" + element = t_meta.pop("element", None) + if element is None: + if "_element_keys" not in t_meta: + raise_error( + msg="`meta` must contain the key 'element' or '_element_keys'" + ) + else: + if isinstance(element, dict): + t_meta["_element_keys"] = list(element.keys()) + else: + t_meta["_element_keys"] = ["element"] + # MD5 hash of the metadata + md5_hash = _meta_hash(t_meta) + return md5_hash, t_meta + + +def element_to_prefix(element: Union[Tuple, Dict, str, int]) -> str: + """Convert the element metadata to prefix. + + Parameters + ---------- + element : tuple, dict, str or int + The element to convert to prefix. + + Returns + ------- + str + The element converted to prefix. + + Raises + ------ + ValueError + If invalid type is passed for `element`. + + """ + logger.debug(f"Converting element {element} to prefix.") + prefix = "element" + if isinstance(element, tuple): + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" + elif isinstance(element, dict): + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" + elif isinstance(element, (str, int)): + prefix = f"{prefix}_{element}" + else: + raise_error( + f"Cannot convert element of type {type(element)} to prefix. " + "Must be a str, int, tuple or dict." + ) + logger.debug(f"Converted prefix: {prefix}") + return f"{prefix}_" + + +def element_to_index( + meta: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None +) -> pd.MultiIndex: + """Convert the element metadata to index. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element"." + n_rows : int, optional + Number of rows to create (default 1). + rows_col_name: str, optional + The column name to use in case `n_rows` > 1. If None and + n_rows > 1, the name will be "idx" (default None). + + Returns + ------- + pandas.MultiIndex + The index of the dataframe to store. + + Raises + ------ + ValueError + If `meta` does not contain the key "element". + + """ + if "element" not in meta: + raise_error( + msg="To create and index, metadata must contain the key 'element'." + ) + # Get element + element = meta["element"] + if not isinstance(element, dict): + element = {"element": element} + # Check rows_col_name + if rows_col_name is None: + rows_col_name = "idx" + elem_idx: Dict[Any, Any] = {k: [v] * n_rows for k, v in element.items()} + elem_idx[rows_col_name] = np.arange(n_rows) + # Create index + index = pd.MultiIndex.from_frame( + pd.DataFrame(elem_idx, index=range(n_rows)) + ) + return index diff --git a/junifer/testing/__init__.py b/junifer/testing/__init__.py new file mode 100644 index 000000000..a1e573655 --- /dev/null +++ b/junifer/testing/__init__.py @@ -0,0 +1,7 @@ +"""Provide imports for testing sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from . import datagrabbers diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py new file mode 100644 index 000000000..9d0f29415 --- /dev/null +++ b/junifer/testing/datagrabbers.py @@ -0,0 +1,67 @@ +"""Provide testing datagrabbers.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import tempfile +from typing import Dict, List + +from nilearn import datasets + +from ..datagrabber.base import BaseDataGrabber + + +class OasisVBMTestingDatagrabber(BaseDataGrabber): + """DataGrabber for Oasis VBM testing data.""" + + def __init__(self) -> None: + """Initialize the class.""" + # Create temporary directory + datadir = tempfile.mkdtemp() + # Define types + types = ["VBM_GM"] + super().__init__(types=types, datadir=datadir) + + def __getitem__(self, element: str) -> Dict: + """Implement indexing support. + + Parameters + ---------- + element : str + The element to retrieve. + + Returns + ------- + dict + The data along with the metadata. + + """ + out = super().__getitem__(element) + i_sub = int(element.split("-")[1]) - 1 + out["VBM_GM"] = {"path": self._dataset.gray_matter_maps[i_sub]} + # Set the element accordingly + out["meta"]["element"] = {"subject": element} + return out + + def __enter__(self) -> "OasisVBMTestingDatagrabber": + """Implement context entry. + + Returns + ------- + OasisVBMTestingDatagrabber + + """ + self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) + return self + + def get_elements(self) -> List[str]: + """Get elements. + + Returns + ------- + list of str + List of elements that can be grabbed. + + """ + return [f"sub-{x:02d}" for x in list(range(1, 11))] diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py new file mode 100644 index 000000000..77024f767 --- /dev/null +++ b/junifer/testing/registry.py @@ -0,0 +1,16 @@ +"""Provide testing registry.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from ..api.registry import register +from .datagrabbers import OasisVBMTestingDatagrabber + + +# Register testing datagrabber +register( + step="datagrabber", + name="OasisVBMTestingDatagrabber", + klass=OasisVBMTestingDatagrabber, +) diff --git a/junifer/tests/test_main.py b/junifer/tests/test_main.py index b5fa71ce7..e29855087 100644 --- a/junifer/tests/test_main.py +++ b/junifer/tests/test_main.py @@ -1,6 +1,13 @@ +"""Provide tests for junifer package.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -def test_import(): + + +def test_import() -> None: + """Test junifer import.""" import junifer + print(junifer.__version__) diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py new file mode 100644 index 000000000..cd1daa60d --- /dev/null +++ b/junifer/tests/test_stats.py @@ -0,0 +1,30 @@ +"""Provide tests for stats.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import Dict, Optional + +import pytest + +from junifer.stats import get_aggfunc_by_name + + +@pytest.mark.parametrize( + "name, params", + [ + ("winsorized_mean", {"limits": [0.2, 0.7]}), + ("mean", None), + ("std", None), + ("trim_mean", None), + ], +) +def test_get_aggfunc_by_name(name: str, params: Optional[Dict]) -> None: + """Test aggregation function retrieval by name.""" + get_aggfunc_by_name(name=name, func_params=params) + + +@pytest.mark.skip(reason="test not implemented") +def test_winsorized_mean() -> None: + """Test winsorized mean computation.""" + ... diff --git a/junifer/utils/__init__.py b/junifer/utils/__init__.py index c8b961c0b..b7fde7ed9 100644 --- a/junifer/utils/__init__.py +++ b/junifer/utils/__init__.py @@ -1,2 +1,8 @@ -from . import logging -from .logging import configure_logging, logger, raise_error, warn \ No newline at end of file +"""Provide imports for utils sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .fs import make_executable +from .logging import configure_logging, logger, raise_error, warn_with_log diff --git a/junifer/utils/fs.py b/junifer/utils/fs.py new file mode 100644 index 000000000..b45e33345 --- /dev/null +++ b/junifer/utils/fs.py @@ -0,0 +1,21 @@ +"""Provide functions for filesystem manipulation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import stat +from pathlib import Path + + +def make_executable(path: Path) -> None: + """Make `path` executable. + + Parameters + ---------- + path : pathlib.Path + The path to make executable. + + """ + st = path.stat() + path.chmod(mode=st.st_mode | stat.S_IEXEC) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index a5558aad8..169a74276 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -1,157 +1,140 @@ +"""Provide class and functions for logging.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL + import logging -import subprocess import sys from distutils.version import LooseVersion from pathlib import Path -import warnings +from subprocess import PIPE, Popen, TimeoutExpired +from typing import Dict, NoReturn, Optional, Type, Union +from warnings import warn -logging.basicConfig(stream=sys.stdout, level=logging.WARN) -logger = logging.getLogger('JUNIFER') +logger = logging.getLogger("JUNIFER") + +_logging_types = { + "DEBUG": logging.DEBUG, + "INFO": logging.INFO, + "WARNING": logging.WARNING, + "ERROR": logging.ERROR, +} -def _get_git_head(path): - """Aux function to read HEAD from git""" +class WrapStdOut(logging.StreamHandler): + """ + Dynamically wrap to sys.stdout. + + This makes packages that monkey-patch sys.stdout (e.g.doctest, + sphinx-gallery) work properly. + + """ + + def __getattr__(self, name: str) -> str: + """Implement attribute fetch.""" + # Even more ridiculous than this class, this must be sys.stdout (not + # just stdout) in order for this to work (tested on OSX and Linux) + if hasattr(sys.stdout, name): + return getattr(sys.stdout, name) + else: + raise AttributeError(f"'file' object has not attribute '{name}'") + + +def _get_git_head(path: Path) -> str: + """Aux function to read HEAD from git. + + Parameters + ---------- + path : pathlib.Path + The path to read git HEAD from. + + Returns + ------- + str + Empty string if timeout expired for subprocess command execution else + git HEAD information. + + """ if not path.exists(): - raise ValueError('This path does not exist: {}'.format(path)) - command = ('cd {gitpath}; ' - 'git rev-parse --verify HEAD').format(gitpath=path) - process = subprocess.Popen(command, - stdout=subprocess.PIPE, - shell=True) - proc_stdout = process.communicate()[0].strip() - del process + raise ValueError(f"This path does not exist: {path}") + command = f"cd {path}; git rev-parse --verify HEAD" + process = Popen( + args=command, + stdout=PIPE, + shell=True, + ) + try: + stdout, _ = process.communicate(timeout=10) + proc_stdout = stdout.strip().decode() + except TimeoutExpired: + process.kill() + proc_stdout = "" return proc_stdout -def get_versions(sys): - """Import stuff and get versions if module - Parameters - ---------- - sys : module - The sys module object. +def get_versions() -> Dict: + """Import stuff and get versions if module. + Returns ------- module_versions : dict The module names and corresponding versions. + """ module_versions = {} for name, module in sys.modules.items(): - if '.' in name: + if "." in name: continue - if name in ['_curses']: + if name in ["_curses"]: continue - vstring = str(getattr(module, '__version__', None)) + vstring = str(getattr(module, "__version__", None)) module_version = LooseVersion(vstring) - module_version = getattr(module_version, 'vstring', None) + module_version = getattr(module_version, "vstring", None) if module_version is None: module_version = None - elif 'git' in module_version: - git_path = Path(module.__file__).resolve().parent + elif "git" in module_version: + git_path = Path(module.__file__).resolve().parent # type: ignore head = _get_git_head(git_path) - module_version += '-HEAD:{}'.format(head) + module_version += f"-HEAD:{head}" module_versions[name] = module_version return module_versions -def get_ext_versions(tbox_path): - """ Get versions of external tools used by JUNIFER.""" - versions = {} - # spm_path = tbox_path / 'spm12' - # if spm_path.exists(): - # head = _get_git_head(spm_path) - # module_version = 'SPM12-HEAD:{}'.format(head) - # versions['spm'] = module_version - return versions +# def get_ext_versions(tbox_path: Path) -> Dict: +# """Get versions of external tools used by junifer. + +# Parameters +# ---------- +# tbox_path : pathlib.Path +# The path to external toolboxes. + +# Returns +# ------- +# dict +# The dependency information. + +# """ +# versions = {} +# # spm_path = tbox_path / 'spm12' +# # if spm_path.exists(): +# # head = _get_git_head(spm_path) +# # module_version = 'SPM12-HEAD:{}'.format(head) +# # versions['spm'] = module_version +# return versions -def _safe_log(versions, name): - if name in versions: - logger.info(f'{name}: {versions[name]}') - - -def log_versions(tbox_path=None): - versions = get_versions(sys) - - logger.info('===== Lib Versions =====') - _safe_log(versions, 'numpy') - _safe_log(versions, 'scipy') - _safe_log(versions, 'pandas') - _safe_log(versions, 'nipype') - _safe_log(versions, 'nitime') - _safe_log(versions, 'nilearn') - _safe_log(versions, 'nibabel') - _safe_log(versions, 'junifer') - logger.info('========================') - - if tbox_path is not None: - # ext_versions = get_ext_versions(tbox_path) - # logger.info('spm: {}'.format(ext_versions['spm'])) - logger.info('========================') - - -_logging_types = dict(DEBUG=logging.DEBUG, INFO=logging.INFO, - WARNING=logging.WARNING, ERROR=logging.ERROR) - - -def configure_logging(level='WARNING', fname=None, overwrite=None, - output_format=None): - """Configure the logging functionality +def _close_handlers(logger: logging.Logger) -> None: + """Safely close relevant handlers for logger. Parameters ---------- - level : int or string - The level of the messages to print. If string, it will be interpreted - as elements of logging. - Options are: ['DEBUG', 'INFO', 'WARNING', 'ERROR']. Defaults to - 'WARNING'. - fname : str, Path or None - Filename of the log to print to. If None, stdout is used. - overwrite : bool | None - Overwrite the log file (if it exists). Otherwise, statements - will be appended to the log (default). None is the same as False, - but additionally raises a warning to notify the user that log - entries will be appended. - output_format : str - Format of the output messages. See the following for examples: + logger : logging.logger + The logger to close handlers for. - https://docs.python.org/dev/howto/logging.html - - e.g., "%(asctime)s - %(levelname)s - %(message)s". - - Defaults to "%(asctime)s - %(name)s - %(levelname)s - %(message)s" """ - _close_handlers(logger) - if output_format is None: - output_format = '%(asctime)s - %(name)s - %(levelname)s - %(message)s' - formatter = logging.Formatter(output_format) - - if fname is not None: - if not isinstance(fname, Path): - fname = Path(fname) - if fname.exists() and overwrite is None: - warnings.warn( - f'File ({fname.as_posix()}) exists. ' - 'Messages will be appended. Use overwrite=True to ' - 'overwrite or overwrite=False to avoid this message') - overwrite = False - mode = 'w' if overwrite else 'a' - lh = logging.FileHandler(fname, mode=mode) - else: - lh = logging.StreamHandler(WrapStdOut()) # type: ignore - - if isinstance(level, str): - level = _logging_types[level] - lh.setFormatter(formatter) - logger.setLevel(level) - logger.addHandler(lh) - log_versions() - - -def _close_handlers(logger): for handler in list(logger.handlers): if isinstance(handler, (logging.FileHandler, logging.StreamHandler)): if isinstance(handler, logging.FileHandler): @@ -159,35 +142,148 @@ def _close_handlers(logger): logger.removeHandler(handler) -def raise_error(msg, klass=ValueError): +def _safe_log(versions: Dict, name: str) -> None: + """Log with safety. + + Parameters + ---------- + versions : dict + The dictionary with keys as dependency names and values as the + versions. + name : str + The dependency to look up in `versions`. + + """ + if name in versions: + logger.info(f"{name}: {versions[name]}") + + +def log_versions(tbox_path: Optional[Path] = None) -> None: + """Log versions of dependencies and junifer. + + If `tbox_path` is specified, can also log versions of external toolboxes. + + Parameters + ---------- + tbox_path : pathlib.Path, optional + The path to external toolboxes (default None). + + """ + # Get versions of all found packages + versions = get_versions() + + logger.info("===== Lib Versions =====") + _safe_log(versions, "numpy") + _safe_log(versions, "scipy") + _safe_log(versions, "pandas") + _safe_log(versions, "nipype") + _safe_log(versions, "nitime") + _safe_log(versions, "nilearn") + _safe_log(versions, "nibabel") + _safe_log(versions, "junifer") + logger.info("========================") + + if tbox_path is not None: + # ext_versions = get_ext_versions(tbox_path) + # logger.info('spm: {}'.format(ext_versions['spm'])) + # logger.info('========================') + pass + + +def configure_logging( + level: Union[int, str] = "WARNING", + fname: Optional[Union[str, Path]] = None, + overwrite: Optional[bool] = None, + output_format=None, +) -> None: + """Configure the logging functionality. + + Parameters + ---------- + level : int or {"DEBUG", "INFO", "WARNING", "ERROR"} + The level of the messages to print. If string, it will be interpreted + as elements of logging (default "WARNING"). + fname : str or pathlib.Path, optional + Filename of the log to print to. If None, stdout is used + (default None). + overwrite : bool, optional + Overwrite the log file (if it exists). Otherwise, statements + will be appended to the log (default). None is the same as False, + but additionally raises a warning to notify the user that log + entries will be appended (default None). + output_format : str, optional + Format of the output messages. See the following for examples: + https://docs.python.org/dev/howto/logging.html + e.g., "%(asctime)s - %(levelname)s - %(message)s". + If None, default string format is used + (default "%(asctime)s - %(name)s - %(levelname)s - %(message)s"). + + """ + _close_handlers(logger) # close relevant logger handlers + + # Set logging level + if isinstance(level, str): + level = _logging_types[level] + + # Set logging output handler + if fname is not None: + # Convert str to Path + if not isinstance(fname, Path): + fname = Path(fname) + if fname.exists() and overwrite is None: + warn( + f"File ({str(fname.absolute())}) exists. " + "Messages will be appended. Use overwrite=True to " + "overwrite or overwrite=False to avoid this message." + ) + overwrite = False + mode = "w" if overwrite else "a" + lh = logging.FileHandler(fname, mode=mode) + else: + lh = logging.StreamHandler(WrapStdOut()) # type: ignore + + # Set logging format + if output_format is None: + output_format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s" + # ( + # "%(asctime)s [%(levelname)s] %(message)s " + # "(%(filename)s:%(lineno)s)" + # ) + formatter = logging.Formatter(fmt=output_format) + + lh.setFormatter(formatter) # set formatter + logger.setLevel(level) # set level + logger.addHandler(lh) # set handler + log_versions() # log versions of installed packages + + +def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: + """Raise error, but first log it. + + Parameters + ---------- + msg : str + The message for the exception. + klass : subclass of Exception, optional + The subclass of Exception to raise using (default ValueError). + + """ logger.error(msg) raise klass(msg) -def warn(msg, category=RuntimeWarning): - """Warn, but first log it +def warn_with_log( + msg: str, category: Optional[Type[Warning]] = RuntimeWarning +) -> None: + """Warn, but first log it. + Parameters ---------- msg : str - Warning message - category : instance of Warning - The warning class. Defaults to ``RuntimeWarning``. + Warning message. + category : subclass of Warning, optional + The warning subclass (default RuntimeWarning). + """ logger.warning(msg) - warnings.warn(msg, category=category) - - -class WrapStdOut(object): - """Dynamically wrap to sys.stdout. - - This makes packages that monkey-patch sys.stdout (e.g.doctest, - sphinx-gallery) work properly. - """ - - def __getattr__(self, name): # noqa: D105 - # Even more ridiculous than this class, this must be sys.stdout (not - # just stdout) in order for this to work (tested on OSX and Linux) - if hasattr(sys.stdout, name): - return getattr(sys.stdout, name) - else: - raise AttributeError(f"'file' object has not attribute '{name}'") + warn(msg, category=category) diff --git a/junifer/utils/tests/test_fs.py b/junifer/utils/tests/test_fs.py new file mode 100644 index 000000000..17912a6a8 --- /dev/null +++ b/junifer/utils/tests/test_fs.py @@ -0,0 +1,30 @@ +"""Provide tests for filesystem manipulation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +from junifer.utils.fs import make_executable + + +def test_make_executable(tmp_path: Path) -> None: + """Test making path executable. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + test_file_path = tmp_path / "make_me_executable.txt" + test_file_path.write_bytes(b"umm") + test_file_stat_initial = test_file_path.stat() + # Check initial file mode + assert test_file_stat_initial.st_mode == 33188 + # Make the path executable + make_executable(test_file_path) + test_file_stat_final = test_file_path.stat() + # Check final file mode + assert test_file_stat_final.st_mode == 33252 diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index 1281fb8f7..a815f8658 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -1,133 +1,221 @@ +"""Provide tests for logging.""" + # Authors: Federico Raimondo # Sami Hamdan +# Synchon Mandal # License: AGPL -from junifer.utils import logger, configure_logging, raise_error, warn -from junifer.utils.logging import _close_handlers -import pytest -import tempfile + +import logging from pathlib import Path +import pytest -def test_log_file(): - """Test logging to a file""" - with tempfile.TemporaryDirectory() as tmp: - tmpdir = Path(tmp) - configure_logging(fname=tmpdir / 'test1.log') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test1.log') as f: +from junifer.utils.logging import ( + _close_handlers, + configure_logging, + get_versions, + log_versions, + logger, + raise_error, + warn_with_log, +) + + +def test_get_versions() -> None: + """Test version info fetch for modules.""" + module_versions = get_versions() + assert "junifer" in module_versions.keys() + + +def test_log_versions(caplog: pytest.LogCaptureFixture) -> None: + """Test logging of dependency and junifer versions. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + # Set log capturing at INFO + with caplog.at_level(logging.INFO): + # Log versions + log_versions() + # Check logging levels + for record in caplog.records: + assert record.levelname not in ("DEBUG", "WARNING", "ERROR") + assert record.levelname == "INFO" + # Check logging message + assert "junifer" in caplog.text + + +def test_log_file(tmp_path: Path) -> None: + """Test logging to a file. + + tmp_path : Path + The path to the test directory. + + """ + configure_logging(fname=tmp_path / "test1.log") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test1.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test2.log", level="INFO") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test2.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test3.log", level="WARNING") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test3.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test4.log", level="ERROR") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + with pytest.warns(UserWarning, match="to avoid this message"): + configure_logging(fname=tmp_path / "test4.log", level="WARNING") + logger.debug("Debug2 message") + logger.info("Info2 message") + logger.warning("Warn2 message") + logger.error("Error2 message") + with open(tmp_path / "test4.log") as f: lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + assert not any("Debug2 message" in line for line in lines) + assert not any("Info2 message" in line for line in lines) + assert any("Warn2 message" in line for line in lines) + assert any("Error2 message" in line for line in lines) - configure_logging(fname=tmpdir / 'test2.log', level='INFO') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test2.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) + configure_logging( + fname=tmp_path / "test4.log", level="WARNING", overwrite=True + ) + logger.debug("Debug3 message") + logger.info("Info3 message") + logger.warning("Warn3 message") + logger.error("Error3 message") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert not any("Error message" in line for line in lines) + assert not any("Debug2 message" in line for line in lines) + assert not any("Info2 message" in line for line in lines) + assert not any("Warn2 message" in line for line in lines) + assert not any("Error2 message" in line for line in lines) + assert not any("Debug3 message" in line for line in lines) + assert not any("Info3 message" in line for line in lines) + assert any("Warn3 message" in line for line in lines) + assert any("Error3 message" in line for line in lines) - configure_logging(fname=tmpdir / 'test3.log', level='WARNING') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test3.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - - configure_logging(fname=tmpdir / 'test4.log', level='ERROR') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - - with pytest.warns(UserWarning, match='to avoid this message'): - configure_logging(fname=tmpdir / 'test4.log', level='WARNING') - logger.debug('Debug2 message') - logger.info('Info2 message') - logger.warning('Warn2 message') - logger.error('Error2 message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - assert not any('Debug2 message' in line for line in lines) - assert not any('Info2 message' in line for line in lines) - assert any('Warn2 message' in line for line in lines) - assert any('Error2 message' in line for line in lines) - - configure_logging(fname=tmpdir / 'test4.log', level='WARNING', - overwrite=True) - logger.debug('Debug3 message') - logger.info('Info3 message') - logger.warning('Warn3 message') - logger.error('Error3 message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert not any('Error message' in line for line in lines) - assert not any('Debug2 message' in line for line in lines) - assert not any('Info2 message' in line for line in lines) - assert not any('Warn2 message' in line for line in lines) - assert not any('Error2 message' in line for line in lines) - assert not any('Debug3 message' in line for line in lines) - assert not any('Info3 message' in line for line in lines) - assert any('Warn3 message' in line for line in lines) - assert any('Error3 message' in line for line in lines) - - with pytest.warns(RuntimeWarning, match=r"Warn raised"): - warn('Warn raised') - with pytest.raises(ValueError, match=r"Error raised"): - raise_error('Error raised') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert any('Warn raised' in line for line in lines) - assert any('Error raised' in line for line in lines) + with pytest.warns(RuntimeWarning, match=r"Warn raised"): + warn_with_log("Warn raised") + with pytest.raises(ValueError, match=r"Error raised"): + raise_error("Error raised") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert any("Warn raised" in line for line in lines) + assert any("Error raised" in line for line in lines) -def test_log(): - """Simple log test""" +def test_log_stdout(caplog: pytest.LogCaptureFixture) -> None: + """Test logging to stdout. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ configure_logging() - logger.info('Testing') + logger.info("Testing") + for record in caplog.records: + assert record.levelname == "INFO" -def test_lib_logging(): - """Test logging versions""" +def test_lib_logging(tmp_path: Path) -> None: + """Test logging versions. + tmp_path : Path + The path to the test directory. + + """ + # Import third-party packages import numpy as np # noqa import pandas # noqa - with tempfile.TemporaryDirectory() as tmp: - tmpdir = Path(tmp) - configure_logging(fname=tmpdir / 'test1.log', level='INFO') - logger.info('first message') - with open(tmpdir / 'test1.log') as f: - lines = f.readlines() - assert any('numpy' in line for line in lines) - assert any('pandas' in line for line in lines) - assert any('junifer' in line for line in lines) + + log_file_path = tmp_path / "test_lib_logging.log" + configure_logging(fname=log_file_path, level="INFO") + logger.info("first message") + with open(log_file_path) as f: + lines = f.readlines() + assert any("numpy" in line for line in lines) + assert any("pandas" in line for line in lines) + assert any("junifer" in line for line in lines) + + +def test_raise_error(caplog: pytest.LogCaptureFixture) -> None: + """Test logging and raising error. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + with pytest.raises(ValueError, match="test error"): + raise_error(msg="test error") + for record in caplog.records: + assert record.levelname == "ERROR" + + +def test_warn_with_log(caplog: pytest.LogCaptureFixture) -> None: + """Test logging and warning. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + with pytest.warns(RuntimeWarning, match="test warning"): + warn_with_log("test warning") + for record in caplog.records: + assert record.levelname == "WARNING" diff --git a/pyproject.toml b/pyproject.toml index 9cae268d7..38d20190e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,85 @@ [build-system] -requires = ["setuptools>=45", "wheel", "setuptools_scm[toml]>=6.2"] +requires = [ + "setuptools >= 61.0.0", + "wheel", + "setuptools_scm[toml] >= 6.2" +] build-backend = "setuptools.build_meta" +[project] +name = "junifer" +description = "JUelich NeuroImaging FEature extractoR" +readme = "README.md" +requires-python = ">=3.8" +license = {file = "LICENSE.md"} +authors = [ + {email = "f.raimondo@fz-juelich.de"}, + {name = "Fede Raimondo"} +] +maintainers = [ + {email = "s.mandal@fz-juelich.de"}, + {name = "Synchon Mandal"} +] +keywords = [ + "neuroimaging", +] +classifiers = [ + "Development Status :: 4 - Beta", + "Intended Audience :: Science/Research", + "Intended Audience :: Developers", + "License :: OSI Approved", + "Natural Language :: English", + "Topic :: Software Development", + "Topic :: Scientific/Engineering", + "Topic :: Scientific/Engineering :: Bio-Informatics", + "Operating System :: OS Independent", + "Programming Language :: Python :: 3.8", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", +] +dependencies = [ + "click>=8.1.3,<8.2", + "numpy>=1.22,<1.23", + "datalad>=0.15.4,<0.18", + "pandas>=1.4.0,<1.5", + "nibabel>=3.2.0,<4.1", + "nilearn>=0.9.0,<1.0", + "sqlalchemy>=1.4.27,<= 1.5.0", + "pyyaml>=5.1.2,<7.0", +] +dynamic = ["version"] + +[project.urls] +homepage = "https://juaml.github.io/junifer" +documentation = "https://juaml.github.io/junifer" +repository = "https://github.com/juaml/junifer" + +[project.scripts] +junifer = "junifer.api.cli:cli" + +[project.optional-dependencies] +dev = ["tox"] +docs = [ + "seaborn>=0.11.2,<0.12", + "Sphinx>=5.0.2,<5.1", + "sphinx-gallery>=0.10.1,<0.11", + "sphinx-rtd-theme>=1.0.0,<1.1", + "sphinx-multiversion>=0.2.4,<0.3", + "numpydoc>=1.4.0,<1.5", +] + +################ +# Tool configs # +################ + +[tool.setuptools] +packages = ["junifer"] + [tool.setuptools_scm] version_scheme = "python-simplified-semver" local_scheme = "no-local-version" write_to = "junifer/_version.py" -write_to_template = "__version__ = '{version}'\n" \ No newline at end of file + +[tool.black] +line-length = 79 +target-version = ["py38"] \ No newline at end of file diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index ea87f8453..000000000 --- a/requirements.txt +++ /dev/null @@ -1,5 +0,0 @@ -numpy>=1.20, <1.22 -datalad>=0.15.4, <0.16 -pandas>=0.18.0, <1.5 -nibabel>=3.2.0, <4.0 -nilearn>=0.9.0, <1.0 \ No newline at end of file diff --git a/scratch/test_niftimasker.py b/scratch/test_niftimasker.py new file mode 100644 index 000000000..db7fcd0a2 --- /dev/null +++ b/scratch/test_niftimasker.py @@ -0,0 +1,76 @@ +from distutils.command.config import config +import numpy as np +from nilearn import datasets +import nibabel as nib +from junifer.utils import configure_logging + +from nilearn.image import resample_to_img, math_img +from nilearn.maskers import NiftiMasker, NiftiLabelsMasker + +configure_logging(level='INFO') + +oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) + +vbm = oasis_dataset.gray_matter_maps[0] +atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + +nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) +auto = nifti_masker.fit_transform(vbm) + + +img = nib.load(vbm) + +atlas_img_res = resample_to_img( + atlas.maps, + img, + interpolation='nearest', +) +atlas_bin = math_img( + 'img != 0', + img=atlas_img_res, +) + + +masker = NiftiMasker(atlas_bin, target_affine=img.affine) + +data = masker.fit_transform(img) +atlas_values = masker.transform(atlas_img_res) +atlas_values = np.squeeze(atlas_values).astype(int) + +manual = [] +for t_v in sorted(np.unique(atlas_values)): + t_values = np.mean(data[:, atlas_values == t_v]) + manual.append(t_values) + + +from junifer.markers.parcel import ParcelAggregation + +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(VBM_GM=dict(data=img)) +jun_values3d = marker.fit_transform(input)['VBM_GM']['data'] + + +# Now do the 4D case +from nilearn.datasets import fetch_spm_auditory +from nilearn.image import concat_imgs +subject_data = fetch_spm_auditory() +fmri_img = concat_imgs(subject_data.func) + +nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) +auto4d = nifti_masker.fit_transform(fmri_img) + +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(BOLD=dict(data=fmri_img)) +jun_values4d = marker.fit_transform(input)['BOLD']['data'] + + +print(auto.shape) +print(jun_values3d.shape) +print(auto4d.shape) +print(jun_values4d.shape) + + +# Now both +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(BOLD=dict(data=fmri_img), VBM_GM=dict(data=img)) +jun_both = marker.fit_transform(input) \ No newline at end of file diff --git a/setup.py b/setup.py index 846a52937..e0709284e 100644 --- a/setup.py +++ b/setup.py @@ -1,59 +1,12 @@ +"""Set up junifer package.""" + # Authors: Federico Raimondo # Sami Hamdan +# Synchon Mandal # License: AGPL -import setuptools -with open('README.md', 'r') as fh: - long_description = fh.read() +from setuptools import setup -def _getversion(): - from setuptools_scm.version import get_local_node_and_date, \ - simplified_semver_version - - def clean_scheme(version): - print(version) - return get_local_node_and_date(version) if version.dirty else "" - - return { - 'version_scheme': simplified_semver_version, - 'local_scheme': clean_scheme, - 'write_to': 'junifer/_version.py', - 'write_to_template': "__version__ = '{version}'\n"} - - -DOWNLOAD_URL = 'https://github.com/juaml/junifer' -URL = 'https://juaml.github.io/junifer' - -setuptools.setup( - name='junifer', - author='Fede Raimondo', - author_email='f.raimondo@fz-juelich.de', - description='JUelich NeuroImaging FEature extractoR', - long_description=long_description, - long_description_content_type='text/markdown', - url=URL, - download_url=DOWNLOAD_URL, - packages=setuptools.find_packages(), - zip_safe=False, - classifiers=['Intended Audience :: Science/Research', - 'Intended Audience :: Developers', - 'License :: OSI Approved', - 'Programming Language :: Python', - 'Topic :: Software Development', - 'Topic :: Scientific/Engineering', - 'Operating System :: Microsoft :: Windows', - 'Operating System :: POSIX', - 'Operating System :: Unix', - 'Operating System :: MacOS', - 'Programming Language :: Python :: 3'], - project_urls={ - 'Documentation': URL, - 'Source': DOWNLOAD_URL, - 'Tracker': f'{DOWNLOAD_URL}issues/', - }, - install_requires=[], # TODO: Complete - python_requires='>=3.6', - use_scm_version=_getversion, - setup_requires=['setuptools_scm'], -) +if __name__ == "__main__": + setup() diff --git a/test-requirements.txt b/test-requirements.txt deleted file mode 100644 index 17b606f42..000000000 --- a/test-requirements.txt +++ /dev/null @@ -1,5 +0,0 @@ -flake8 -pytest -pytest-cov -codecov -https://github.com/codespell-project/codespell/archive/master.zip \ No newline at end of file diff --git a/tools/create_bids_example_dataset_sessions.py b/tools/create_bids_example_dataset_sessions.py new file mode 100644 index 000000000..3301c6f8a --- /dev/null +++ b/tools/create_bids_example_dataset_sessions.py @@ -0,0 +1,41 @@ +# Authors: Federico Raimondo +# License: AGPL +from tempfile import TemporaryDirectory +from pathlib import Path + +import datalad.api as dl + +dst = 'git@gin.g-node.org:/juaml/datalad-example-bids-ses.git' + +with TemporaryDirectory() as tmpdir_name: + tmpdir = Path(tmpdir_name) + ds = dl.create(tmpdir) # type: ignore + + base_dir = tmpdir / 'example_bids_ses' + base_dir.mkdir() + + for i_sub in range(1, 10): + t_sub = f'sub-{i_sub:02d}' + sub_dir = base_dir / t_sub + sub_dir.mkdir() + + for i_ses in range(1, 4): + t_ses = f'ses-{i_ses:02d}' + ses_dir = sub_dir / t_ses + ses_dir.mkdir() + + for dname in ['anat', 'func']: + (ses_dir / dname).mkdir() + + fnames = [f'anat/{t_sub}_{t_ses}_T1w.nii.gz'] + if i_ses != 3: # Session 3 does not have functional data + fnames.extend([ + f'func/{t_sub}_{t_ses}_task-rest_bold.nii.gz', + f'func/{t_sub}_{t_ses}_task-rest_bold.json']) + for fname in fnames: + with open(ses_dir / fname, 'w') as f: + f.write('placeholder') + + ds.save(recursive=True) + ds.siblings('add', name='gin', url=dst) + ds.push(to='gin', force='all') diff --git a/tox.ini b/tox.ini new file mode 100644 index 000000000..10f890c2e --- /dev/null +++ b/tox.ini @@ -0,0 +1,150 @@ +[tox] +envlist = isort, black, flake8, test, coverage, codespell, py3{8,9,10} +isolated_build = true + +[gh-actions] +python = + 3.8: py38 + 3.9: coverage + 3.10: py310 + +[testenv] +skip_install = false +# Required for git-annex +passenv = + HOME +deps = + pytest +commands = + pytest + +[testenv:isort] +skip_install = true +deps = + isort +commands = + isort --check-only --diff {toxinidir}/junifer {toxinidir}/setup.py + +[testenv:black] +skip_install = true +deps = + black +commands = + black --check --diff {toxinidir}/junifer {toxinidir}/setup.py + +[testenv:flake8] +skip_install = true +deps = + flake8 + flake8-docstrings + flake8-bugbear +commands = + flake8 {toxinidir}/junifer {toxinidir}/setup.py + +[testenv:test] +skip_install = false +passenv = + HOME +deps = + pytest +commands = + pytest -vv + +[testenv:coverage] +skip_install = false +deps = + pytest + pytest-cov +commands = + pytest --cov={envsitepackagesdir}/junifer --cov-report=xml -vv --cov-report=term + +[testenv:codespell] +skip_install = true +deps = + codespell +commands = + codespell --config tox.ini examples/ junifer/ scratch/ tools/ + +################ +# Tool configs # +################ + +[isort] +skip = + __init__.py +profile = black +line_length = 79 +lines_after_imports = 2 +known_first_party = junifer +known_third_party = + click + numpy + datalad + pandas + nibabel + nilearn + sqlalchemy + yaml + pytest + +[flake8] +exclude = + __init__.py +max-line-length = 79 +extend-ignore = + B024 # abstract class with no abstract methods + D202 + E201 # whitespace after ‘(’ + E202 # whitespace before ‘)’ + E203 # whitespace before ‘,’, ‘;’, or ‘:’ + E221 # multiple spaces before operator + E222 # multiple spaces after operator + E241 # multiple spaces after ‘,’ + I100 + I101 + I201 + N806 + W503 # line break before binary operator + W504 # line break after binary operator + +[pytest] +testpaths = + junifer/api/tests + junifer/configs/tests + junifer/data/tests + junifer/datagrabber/tests + junifer/datareader/tests + junifer/markers/tests + junifer/preprocess/tests + junifer/storage/tests + junifer/testing + junifer/utils/tests + +[coverage:paths] +source = + junifer + */site-packages/junifer + +[coverage:run] +branch = true +omit = + */setup.py + */_version.py + */tests/* +parallel = false + +[coverage:report] +exclude_lines = + # Have to re-enable the standard pragma + pragma: no cover + # Don't complain if non-runnable code isn't run: + if __name__ == .__main__.: +precision = 2 + +[codespell] +skip = docs/auto_*,*.html,.git/,*.pyc,docs/_build +count = +quiet-level = 3 +ignore-words = ignore_words.txt +interactive = 0 +builtin = clear,rare,informal,names,usage