First working prototype of Junifer #3
121 changed files with 10046 additions and 1428 deletions
|
|
@ -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
|
||||
12
.coveragerc
12
.coveragerc
|
|
@ -1,12 +0,0 @@
|
|||
[run]
|
||||
branch = True
|
||||
source = junifer
|
||||
include = */junifer/*
|
||||
omit =
|
||||
*/setup.py
|
||||
*/tests/*
|
||||
|
||||
[report]
|
||||
exclude_lines =
|
||||
pragma: no cover
|
||||
if __name__ == .__main__.:
|
||||
5
.flake8
5
.flake8
|
|
@ -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
|
||||
56
.github/workflows/ci.yml
vendored
56
.github/workflows/ci.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
110
.github/workflows/docs.yml
vendored
110
.github/workflows/docs.yml
vendored
|
|
@ -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
31
.github/workflows/lint.yml
vendored
Normal 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
|
||||
4
.gitignore
vendored
4
.gitignore
vendored
|
|
@ -132,4 +132,6 @@ cython_debug/
|
|||
# OS Stuff
|
||||
.DS_store
|
||||
|
||||
junifer/_version.py
|
||||
junifer/_version.py
|
||||
scratch/
|
||||
junifer_jobs/
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
Original Authors
|
||||
================
|
||||
* Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||
* Leonard Sasse <l.sasse@fz-juelich.de>
|
||||
* Leonard Sasse <l.sasse@fz-juelich.de>
|
||||
* Synchon Mandal <s.mandal@fz-juelich.de>
|
||||
|
|
|
|||
15
Makefile
15
Makefile
|
|
@ -1,15 +0,0 @@
|
|||
# Makefile before PR
|
||||
#
|
||||
|
||||
.PHONY: checks
|
||||
|
||||
checks: flake spellcheck
|
||||
|
||||
flake:
|
||||
flake8
|
||||
|
||||
spellcheck:
|
||||
codespell junifer/ docs/ examples/
|
||||
|
||||
test:
|
||||
pytest -v
|
||||
62
README.md
62
README.md
|
|
@ -1,20 +1,68 @@
|
|||
# python-library-mockup
|
||||
JUelich NeuroImaging FEature extractoR
|
||||
# junifer - JUelich NeuroImaging FEature extractoR
|
||||
|
||||

|
||||

|
||||

|
||||

|
||||
[](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
32
conda-env.yml
Normal 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
|
||||
|
|
@ -1,2 +0,0 @@
|
|||
flake8
|
||||
pytest
|
||||
|
|
@ -1,6 +0,0 @@
|
|||
seaborn
|
||||
sphinx
|
||||
sphinx-gallery
|
||||
sphinx_rtd_theme
|
||||
git+https://github.com/dls-controls/sphinx-multiversion.git@only-arg
|
||||
numpydoc
|
||||
17
docs/api.rst
17
docs/api.rst
|
|
@ -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
7
docs/api/api.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
API Functions
|
||||
^^^^^^^^^^^^^^
|
||||
|
||||
.. automodule:: junifer.api
|
||||
:members:
|
||||
:imported-members:
|
||||
6
docs/api/datagrabbers.rst
Normal file
6
docs/api/datagrabbers.rst
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
Data Grabbers
|
||||
^^^^^^^^^^^^^
|
||||
|
||||
.. automodule:: junifer.datagrabber
|
||||
:members:
|
||||
:imported-members:
|
||||
7
docs/api/datareaders.rst
Normal file
7
docs/api/datareaders.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
Data Readers
|
||||
^^^^^^^^^^^^
|
||||
|
||||
.. automodule:: junifer.datareader
|
||||
:members:
|
||||
:imported-members:
|
||||
27
docs/api/index.rst
Normal file
27
docs/api/index.rst
Normal 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
7
docs/api/markers.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
Markers
|
||||
^^^^^^^
|
||||
|
||||
.. automodule:: junifer.markers
|
||||
:members:
|
||||
:imported-members:
|
||||
7
docs/api/preprocessing.rst
Normal file
7
docs/api/preprocessing.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
Pre-processing
|
||||
^^^^^^^^^^^^^^
|
||||
|
||||
.. automodule:: junifer.preprocess
|
||||
:members:
|
||||
:imported-members:
|
||||
7
docs/api/storage.rst
Normal file
7
docs/api/storage.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
Storage
|
||||
^^^^^^^
|
||||
|
||||
.. automodule:: junifer.storage
|
||||
:members:
|
||||
:imported-members:
|
||||
6
docs/api/testing.rst
Normal file
6
docs/api/testing.rst
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
|
||||
Testing
|
||||
^^^^^^^
|
||||
|
||||
.. automodule:: junifer.testing.datagrabbers
|
||||
:members:
|
||||
7
docs/api/utils.rst
Normal file
7
docs/api/utils.rst
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
|
||||
Utils
|
||||
^^^^^
|
||||
|
||||
.. automodule:: junifer.utils
|
||||
:members:
|
||||
:imported-members:
|
||||
112
docs/builtin.rst
Normal file
112
docs/builtin.rst
Normal 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` | |
|
||||
+------------------+-----------------------+-----------------------------+---------------+
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
.. _Fede Raimondo: https://fraimondo.github.io
|
||||
.. _Kaustubh Patil: https://github.com/kaurao
|
||||
.. _Leonard Sasse: https://github.com/LeSasse
|
||||
.. _Amir Omidvarnia: https://github.com/omidvarnia
|
||||
.. _Amir Omidvarnia: https://github.com/omidvarnia
|
||||
.. _Synchon Mandal: https://github.com/synchon
|
||||
|
|
@ -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
217
docs/contribution.rst
Normal 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
4
docs/faq.rst
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
.. include:: links.inc
|
||||
|
||||
FAQs
|
||||
====
|
||||
|
|
@ -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`
|
||||
|
||||
|
|
|
|||
|
|
@ -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>`_.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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], []
|
||||
|
|
|
|||
46
docs/understanding/data.rst
Normal file
46
docs/understanding/data.rst
Normal 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)
|
||||
6
docs/understanding/datagrabber.rst
Normal file
6
docs/understanding/datagrabber.rst
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
.. include:: ../links.inc
|
||||
|
||||
.. _datagrabber:
|
||||
|
||||
Data Grabber
|
||||
============
|
||||
6
docs/understanding/datareader.rst
Normal file
6
docs/understanding/datareader.rst
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
.. include:: ../links.inc
|
||||
|
||||
.. _datareader:
|
||||
|
||||
Data Reader
|
||||
===========
|
||||
28
docs/understanding/index.rst
Normal file
28
docs/understanding/index.rst
Normal 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
|
||||
4
docs/understanding/marker.rst
Normal file
4
docs/understanding/marker.rst
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
.. include:: ../links.inc
|
||||
|
||||
Marker
|
||||
======
|
||||
4
docs/understanding/storage.rst
Normal file
4
docs/understanding/storage.rst
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
.. include:: ../links.inc
|
||||
|
||||
Storage
|
||||
=======
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
57
examples/run_compute_parcel_mean.py
Normal file
57
examples/run_compute_parcel_mean.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
53
examples/run_run_gmd_mean.py
Normal file
53
examples/run_run_gmd_mean.py
Normal 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,
|
||||
)
|
||||
25
examples/yamls/gmd_mean.yaml
Normal file
25
examples/yamls/gmd_mean.yaml
Normal 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
|
||||
|
||||
20
examples/yamls/gmd_mean_htcondor.yaml
Normal file
20
examples/yamls/gmd_mean_htcondor.yaml
Normal 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
|
||||
25
examples/yamls/ukb_gmd_mean.yaml
Normal file
25
examples/yamls/ukb_gmd_mean.yaml
Normal 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
|
||||
|
||||
|
|
@ -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 configs
|
||||
from . import data
|
||||
from . import datagrabber
|
||||
from . import datareader
|
||||
from . import markers
|
||||
from . import configs
|
||||
from . import pipeline
|
||||
from . import preprocess
|
||||
from . import storage
|
||||
from . import utils
|
||||
|
|
|
|||
|
|
@ -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
195
junifer/api/cli.py
Normal 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
|
||||
|
|
@ -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
571
junifer/api/functions.py
Normal 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
51
junifer/api/parser.py
Normal 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
|
||||
|
|
@ -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
146
junifer/api/registry.py
Normal 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
17
junifer/api/res/run_conda.sh
Executable 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"
|
||||
"$@"
|
||||
15
junifer/api/tests/data/gmd_mean.yaml
Normal file
15
junifer/api/tests/data/gmd_mean.yaml
Normal 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
|
||||
|
||||
20
junifer/api/tests/data/gmd_mean_htcondor.yaml
Normal file
20
junifer/api/tests/data/gmd_mean_htcondor.yaml
Normal 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
|
||||
70
junifer/api/tests/test_cli.py
Normal file
70
junifer/api/tests/test_cli.py
Normal 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
|
||||
157
junifer/api/tests/test_functions.py
Normal file
157
junifer/api/tests/test_functions.py
Normal 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
|
||||
77
junifer/api/tests/test_parser.py
Normal file
77
junifer/api/tests/test_parser.py
Normal 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)
|
||||
136
junifer/api/tests/test_registry.py
Normal file
136
junifer/api/tests/test_registry.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -1 +1,7 @@
|
|||
from .atlases import list_atlases, register_atlas, load_atlas
|
||||
"""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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
483
junifer/data/tests/test_atlases.py
Normal file
483
junifer/data/tests/test_atlases.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
167
junifer/datagrabber/datalad_base.py
Normal file
167
junifer/datagrabber/datalad_base.py
Normal 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
201
junifer/datagrabber/hcp.py
Normal 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,
|
||||
)
|
||||
111
junifer/datagrabber/multiple.py
Normal file
111
junifer/datagrabber/multiple.py
Normal 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
|
||||
204
junifer/datagrabber/pattern.py
Normal file
204
junifer/datagrabber/pattern.py
Normal 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)
|
||||
53
junifer/datagrabber/pattern_datalad.py
Normal file
53
junifer/datagrabber/pattern_datalad.py
Normal 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
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
22
junifer/datagrabber/tests/test_datalad_base.py
Normal file
22
junifer/datagrabber/tests/test_datalad_base.py
Normal 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(
|
||||
|
||||
# )
|
||||
94
junifer/datagrabber/tests/test_multiple.py
Normal file
94
junifer/datagrabber/tests/test_multiple.py
Normal 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)
|
||||
110
junifer/datagrabber/tests/test_pattern.py
Normal file
110
junifer/datagrabber/tests/test_pattern.py
Normal 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"]
|
||||
178
junifer/datagrabber/tests/test_pattern_datalad.py
Normal file
178
junifer/datagrabber/tests/test_pattern_datalad.py
Normal 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)
|
||||
77
junifer/datagrabber/utils.py
Normal file
77
junifer/datagrabber/utils.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
from .default import DefaultDataReader
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# License: AGPL
|
||||
|
||||
from .base import BaseMarker
|
||||
from .collection import MarkerCollection
|
||||
from .parcel import ParcelAggregation
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
144
junifer/markers/parcel.py
Normal 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
|
||||
165
junifer/markers/tests/test_collection.py
Normal file
165
junifer/markers/tests/test_collection.py
Normal 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
|
||||
78
junifer/markers/tests/test_markers_base.py
Normal file
78
junifer/markers/tests/test_markers_base.py
Normal 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"
|
||||
166
junifer/markers/tests/test_parcel.py
Normal file
166
junifer/markers/tests/test_parcel.py
Normal 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"] == {}
|
||||
6
junifer/pipeline/__init__.py
Normal file
6
junifer/pipeline/__init__.py
Normal 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
|
||||
107
junifer/pipeline/pipeline_mixin.py
Normal file
107
junifer/pipeline/pipeline_mixin.py
Normal 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,
|
||||
)
|
||||
27
junifer/pipeline/tests/test_pipeline_mixin.py
Normal file
27
junifer/pipeline/tests/test_pipeline_mixin.py
Normal 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"
|
||||
|
|
@ -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
|
||||
# License: AGPL
|
||||
|
||||
from .confounds import BaseConfoundRemover
|
||||
|
|
|
|||
487
junifer/preprocess/confounds.py
Normal file
487
junifer/preprocess/confounds.py
Normal 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')
|
||||
347
junifer/preprocess/tests/test_confounds.py
Normal file
347
junifer/preprocess/tests/test_confounds.py
Normal 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
111
junifer/stats.py
Normal 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
|
||||
|
|
@ -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>
|
||||
# License: AGPL
|
||||
# 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
263
junifer/storage/base.py
Normal 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}>"
|
||||
65
junifer/storage/pandas_base.py
Normal file
65
junifer/storage/pandas_base.py
Normal 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
598
junifer/storage/sqlite.py
Normal 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
Loading…
Reference in a new issue