First working prototype of Junifer #3

Merged
fraimondo merged 293 commits from dev into main 2022-09-12 19:42:05 +00:00
121 changed files with 10046 additions and 1428 deletions

View file

@ -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

View file

@ -1,12 +0,0 @@
[run]
branch = True
source = junifer
include = */junifer/*
omit =
*/setup.py
*/tests/*
[report]
exclude_lines =
pragma: no cover
if __name__ == .__main__.:

View file

@ -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

View file

@ -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

View file

@ -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

31
.github/workflows/lint.yml vendored Normal file
View file

@ -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

2
.gitignore vendored
View file

@ -133,3 +133,5 @@ cython_debug/
.DS_store
junifer/_version.py
scratch/
junifer_jobs/

View file

@ -2,3 +2,4 @@ Original Authors
================
* Federico Raimondo <f.raimondo@fz-juelich.de>
* Leonard Sasse <l.sasse@fz-juelich.de>
* Synchon Mandal <s.mandal@fz-juelich.de>

View file

@ -1,15 +0,0 @@
# Makefile before PR
#
.PHONY: checks
checks: flake spellcheck
flake:
flake8
spellcheck:
codespell junifer/ docs/ examples/
test:
pytest -v

View file

@ -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)
## 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 <http://www.gnu.org/licenses/>.

32
conda-env.yml Normal file
View file

@ -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

View file

@ -1,2 +0,0 @@
flake8
pytest

View file

@ -1,6 +0,0 @@
seaborn
sphinx
sphinx-gallery
sphinx_rtd_theme
git+https://github.com/dls-controls/sphinx-multiversion.git@only-arg
numpydoc

View file

@ -1,17 +0,0 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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:

7
docs/api/api.rst Normal file
View file

@ -0,0 +1,7 @@
API Functions
^^^^^^^^^^^^^^
.. automodule:: junifer.api
:members:
:imported-members:

View file

@ -0,0 +1,6 @@
Data Grabbers
^^^^^^^^^^^^^
.. automodule:: junifer.datagrabber
:members:
:imported-members:

7
docs/api/datareaders.rst Normal file
View file

@ -0,0 +1,7 @@
Data Readers
^^^^^^^^^^^^
.. automodule:: junifer.datareader
:members:
:imported-members:

27
docs/api/index.rst Normal file
View file

@ -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

7
docs/api/markers.rst Normal file
View file

@ -0,0 +1,7 @@
Markers
^^^^^^^
.. automodule:: junifer.markers
:members:
:imported-members:

View file

@ -0,0 +1,7 @@
Pre-processing
^^^^^^^^^^^^^^
.. automodule:: junifer.preprocess
:members:
:imported-members:

7
docs/api/storage.rst Normal file
View file

@ -0,0 +1,7 @@
Storage
^^^^^^^
.. automodule:: junifer.storage
:members:
:imported-members:

6
docs/api/testing.rst Normal file
View file

@ -0,0 +1,6 @@
Testing
^^^^^^^
.. automodule:: junifer.testing.datagrabbers
:members:

7
docs/api/utils.rst Normal file
View file

@ -0,0 +1,7 @@
Utils
^^^^^
.. automodule:: junifer.utils
:members:
:imported-members:

112
docs/builtin.rst Normal file
View file

@ -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:`<issue number>`
.. list-table:: Available data grabbers
:widths: auto
:header-rows: 1
* - Class
- Description
- Access
- Type/Config
- State
- Version Added
* - `DataladHCP1200`
- `HCP OpenAccess dataset <https://github.com/datalad-datasets/human-connectome-project-openaccess>`_
- 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:`<issue number>`
.. 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` | |
+------------------+-----------------------+-----------------------------+---------------+

View file

@ -2,3 +2,4 @@
.. _Kaustubh Patil: https://github.com/kaurao
.. _Leonard Sasse: https://github.com/LeSasse
.. _Amir Omidvarnia: https://github.com/omidvarnia
.. _Synchon Mandal: https://github.com/synchon

View file

@ -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)

217
docs/contribution.rst Normal file
View file

@ -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
<https://guides.github.com/activities/forking/>`_.
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
<https://realpython.com/python-virtual-environments-a-primer/>`_ 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 <prefix>/<name-of-your-branch>
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
<http://karma-runner.github.io/2.0/dev/git-commit-msg.html>`_.
.. code-block:: console
git add .
git commit -m "<prefix>: <summary of changes>"
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 <prefix>/<name-of-your-branch>
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 <http://docs.pytest.org/en/latest/>`_ 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.

4
docs/faq.rst Normal file
View file

@ -0,0 +1,4 @@
.. include:: links.inc
FAQs
====

View file

@ -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`

View file

@ -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 <https://realpython.com/python-virtual-environments-a-primer>`_.
.. _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 <https://pypi.org>`_, 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 <contribution.rst>`_.

View file

@ -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
.. _`sphinx reST reference`: https://www.sphinx-doc.org/en/master/usage/restructuredtext/basics.html#inline-markup

View file

@ -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], []

View file

@ -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)

View file

@ -0,0 +1,6 @@
.. include:: ../links.inc
.. _datagrabber:
Data Grabber
============

View file

@ -0,0 +1,6 @@
.. include:: ../links.inc
.. _datareader:
Data Reader
===========

View file

@ -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

View file

@ -0,0 +1,4 @@
.. include:: ../links.inc
Marker
======

View file

@ -0,0 +1,4 @@
.. include:: ../links.inc
Storage
=======

View file

@ -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,
)

View file

@ -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"},
)

View file

@ -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)

View file

@ -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)

View file

@ -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,
)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -1,6 +1,17 @@
from . _version import __version__
"""Provide imports for junifer package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from ._version import __version__
from . import api
from . import utils
from . import datagrabber
from . import markers
from . import configs
from . import data
from . import datagrabber
from . import datareader
from . import markers
from . import pipeline
from . import preprocess
from . import storage
from . import utils

View file

@ -1 +1,8 @@
from . pipeline import run_pipeline
"""Provide imports for api sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .cli import cli
from .functions import run, collect

195
junifer/api/cli.py Normal file
View file

@ -0,0 +1,195 @@
"""Provide functions for cli."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,13 +1,17 @@
"""Provide decorators for api."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

571
junifer/api/functions.py Normal file
View file

@ -0,0 +1,571 @@
"""Provide functions for cli."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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}`"
# )

51
junifer/api/parser.py Normal file
View file

@ -0,0 +1,51 @@
"""Provide functions for parser."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,60 +0,0 @@
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# 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 = {}

146
junifer/api/registry.py Normal file
View file

@ -0,0 +1,146 @@
"""Provide functions for registry."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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_

17
junifer/api/res/run_conda.sh Executable file
View file

@ -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"
"$@"

View file

@ -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

View file

@ -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

View file

@ -0,0 +1,70 @@
"""Provide tests for cli."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,157 @@
"""Provide tests for functions."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,77 @@
"""Provide tests for parser."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -0,0 +1,136 @@
"""Provide tests for registry."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -1,65 +1,44 @@
from ..datagrabber import DataladDataGrabber
"""Provide class for juseless datalad datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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,
)

View file

@ -1,15 +1,47 @@
"""Provide tests for juseless datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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()

View file

@ -1 +1,7 @@
"""Provide imports for data sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .atlases import list_atlases, register_atlas, load_atlas

View file

@ -1,17 +1,30 @@
"""Provide functions for atlases."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Vera Komeyer <v.komeyer@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,483 @@
"""Provide tests for atlas."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Vera Komeyer <v.komeyer@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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",
)

View file

@ -1,4 +1,12 @@
"""Provide imports for datagrabber sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# License: AGPL
from .base import BIDSDataladDataGrabber, DataladDataGrabber, BIDSDataGrabber
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

View file

@ -1,294 +1,158 @@
"""Provide abstract base class for datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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,
)

View file

@ -0,0 +1,167 @@
"""Provide abstract base class for datalad datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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")

201
junifer/datagrabber/hcp.py Normal file
View file

@ -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,
)

View file

@ -0,0 +1,111 @@
"""Provide abstract base class for multiple source datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,204 @@
"""Provide concrete implementation for pattern-based datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -0,0 +1,53 @@
"""Provide base class for pattern-based datalad datagrabber."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,76 +1,43 @@
"""Provide tests for base."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"]

View file

@ -0,0 +1,22 @@
"""Provide tests for datalad_base."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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(
# )

View file

@ -0,0 +1,94 @@
"""Provide tests for multiple."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# 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)

View file

@ -0,0 +1,110 @@
"""Provide tests for pattern."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"]

View file

@ -0,0 +1,178 @@
"""Provide tests for pattern_datalad."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -0,0 +1,77 @@
"""Provide utility functions for the datagrabber sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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
)

View file

@ -1,5 +1,8 @@
"""Provide imports for datareader sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .default import DefaultDataReader

View file

@ -1,66 +1,118 @@
"""Provide class for default data reader."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,128 +1,168 @@
"""Provide tests for default data reader."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

View file

@ -1,3 +1,9 @@
"""Provide imports for markers sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# License: AGPL
from .base import BaseMarker
from .collection import MarkerCollection
from .parcel import ParcelAggregation

View file

@ -1,62 +1,155 @@
"""Provide base class for markers."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,70 +1,111 @@
"""Provide class for marker collection."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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)

144
junifer/markers/parcel.py Normal file
View file

@ -0,0 +1,144 @@
"""Provide class for parcel aggregation."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,165 @@
"""Provide tests for marker collection."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -0,0 +1,78 @@
"""Provide tests for base marker."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"

View file

@ -0,0 +1,166 @@
"""Provide test for parcel aggregation."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"] == {}

View file

@ -0,0 +1,6 @@
"""Provide imports for pipeline sub-package."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .pipeline_mixin import PipelineStepMixin

View file

@ -0,0 +1,107 @@
"""Provide mixin class for pipeline step."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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,
)

View file

@ -0,0 +1,27 @@
"""Provide tests for pipeline mixin."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"

View file

@ -1,3 +1,7 @@
"""Provide imports for preprocess sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# License: AGPL
from .confounds import BaseConfoundRemover

View file

@ -0,0 +1,487 @@
"""Provide base class for confound removal."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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')

View file

@ -0,0 +1,347 @@
"""Provide tests for confound removal."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

111
junifer/stats.py Normal file
View file

@ -0,0 +1,111 @@
"""Provide functions for statistics."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -1,3 +1,9 @@
"""Provide imports for storage sub-package."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from .base import BaseFeatureStorage
from .pandas_base import PandasBaseFeatureStorage
from .sqlite import SQLiteFeatureStorage

263
junifer/storage/base.py Normal file
View file

@ -0,0 +1,263 @@
"""Provide abstract base class for feature storage."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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}>"

View file

@ -0,0 +1,65 @@
"""Provide abstract base class for feature storage via pandas."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

598
junifer/storage/sqlite.py Normal file
View file

@ -0,0 +1,598 @@
"""Provide concrete implementation for feature storage via SQLite."""
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

Some files were not shown because too many files have changed in this diff Show more