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:
|
jobs:
|
||||||
build:
|
build:
|
||||||
|
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
python-version: [3.6, 3.7, 3.8]
|
python-version: ['3.8', '3.9', '3.10']
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v2
|
- name: Set up system
|
||||||
- name: Set up Python ${{ matrix.python-version }}
|
|
||||||
uses: actions/setup-python@v2
|
|
||||||
with:
|
|
||||||
python-version: ${{ matrix.python-version }}
|
|
||||||
- name: Check for sudo
|
|
||||||
shell: bash
|
|
||||||
run: |
|
run: |
|
||||||
if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi
|
bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)"
|
||||||
echo "SUDO=$SUDO" >> $GITHUB_ENV
|
sudo apt-get update -qq
|
||||||
- name: Install dependencies
|
sudo apt-get install git-annex-standalone
|
||||||
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
|
|
||||||
- name: Configure git for datalad
|
- name: Configure git for datalad
|
||||||
run: |
|
run: |
|
||||||
git config --global user.email "runner@github.com"
|
git config --global user.email "runner@github.com"
|
||||||
git config --global user.name "GITHUB CI Runner"
|
git config --global user.name "GitHub Runner"
|
||||||
- name: Install junifer
|
- uses: actions/checkout@v3
|
||||||
shell: bash -el {0}
|
- name: Set up Python ${{ matrix.python-version }}
|
||||||
|
uses: actions/setup-python@v4
|
||||||
|
with:
|
||||||
|
python-version: ${{ matrix.python-version }}
|
||||||
|
- name: Install dependencies
|
||||||
run: |
|
run: |
|
||||||
python setup.py build
|
python -m pip install --upgrade pip setuptools wheel
|
||||||
python setup.py install
|
python -m pip install tox tox-gh-actions
|
||||||
- name: Lint with flake8
|
- name: Test with tox
|
||||||
run: |
|
run: |
|
||||||
# stop the build if there are Python syntax errors or undefined names
|
tox
|
||||||
flake8 . --count --show-source --statistics
|
- name: Upload coverage to Codecov
|
||||||
- name: Spell check
|
uses: codecov/codecov-action@v3
|
||||||
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
|
|
||||||
with:
|
with:
|
||||||
token: ${{ secrets.CODECOV_TOKEN }}
|
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
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
strategy:
|
||||||
fail-fast: false
|
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 }}
|
steps:
|
||||||
uses: actions/setup-python@v2
|
- name: Checkout source
|
||||||
with:
|
uses: actions/checkout@v2
|
||||||
python-version: 3.8
|
with:
|
||||||
- name: Check for sudo
|
# require all of history to see all tagged versions' docs
|
||||||
shell: bash
|
fetch-depth: 0
|
||||||
run: |
|
- name: Set up Python 3.9
|
||||||
if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi
|
uses: actions/setup-python@v2
|
||||||
echo "SUDO=$SUDO" >> $GITHUB_ENV
|
with:
|
||||||
- name: Install Dependencies
|
python-version: 3.9
|
||||||
run: |
|
- name: Check for sudo
|
||||||
$SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)"
|
shell: bash
|
||||||
$SUDO apt-get update -qq
|
run: |
|
||||||
$SUDO apt-get install git-annex-standalone
|
if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi
|
||||||
python -m pip install --upgrade pip
|
echo "SUDO=$SUDO" >> $GITHUB_ENV
|
||||||
pip install -r requirements.txt
|
- name: Install dependencies
|
||||||
pip install -r docs-requirements.txt
|
run: |
|
||||||
python setup.py build
|
$SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)"
|
||||||
python setup.py install
|
$SUDO apt-get update -qq
|
||||||
- name: Configure git for datalad
|
$SUDO apt-get install git-annex-standalone
|
||||||
run: |
|
python -m pip install --upgrade pip setuptools wheel
|
||||||
git config --global user.email "runner@github.com"
|
python -m pip install -e .[docs]
|
||||||
git config --global user.name "GITHUB CI Runner"
|
- name: Configure git for datalad
|
||||||
- name: Checkout gh-pages
|
run: |
|
||||||
# As we already did a deploy of gh-pages above, it is guaranteed to be there
|
git config --global user.email "runner@github.com"
|
||||||
# so check it out so we can selectively build docs below
|
git config --global user.name "GITHUB CI Runner"
|
||||||
uses: actions/checkout@v2
|
- name: Checkout gh-pages
|
||||||
with:
|
# 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
|
ref: gh-pages
|
||||||
path: docs/_build
|
path: docs/_build
|
||||||
|
- name: Test build docs
|
||||||
- name: Test Build Docs
|
if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags')
|
||||||
if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags')
|
run: |
|
||||||
run: |
|
BUILDDIR=_build/main make -C docs/ local
|
||||||
BUILDDIR=_build/main make -C docs/ local
|
- name: Build docs
|
||||||
|
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags')
|
||||||
- name: Build Docs
|
# Use the args we normally pass to sphinx-build, but run sphinx-multiversion
|
||||||
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags')
|
run: |
|
||||||
# Use the args we normally pass to sphinx-build, but run sphinx-multiversion
|
make -C docs/ html
|
||||||
run: |
|
touch docs/_build/.nojekyll
|
||||||
make -C docs/ html
|
cp docs/redirect.html docs/_build/index.html
|
||||||
touch docs/_build/.nojekyll
|
- name: Publish docs to gh-pages
|
||||||
cp docs/redirect.html docs/_build/index.html
|
# Only once from main or a tag
|
||||||
|
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags')
|
||||||
- name: Publish Docs to gh-pages
|
# We pin to the SHA, not the tag, for security reasons.
|
||||||
# Only once from main or a tag
|
# https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions
|
||||||
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags')
|
uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3
|
||||||
# We pin to the SHA, not the tag, for security reasons.
|
with:
|
||||||
# https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions
|
github_token: ${{ secrets.GITHUB_TOKEN }}
|
||||||
uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3
|
publish_dir: docs/_build
|
||||||
with:
|
keep_files: true
|
||||||
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
|
# OS Stuff
|
||||||
.DS_store
|
.DS_store
|
||||||
|
|
||||||
junifer/_version.py
|
junifer/_version.py
|
||||||
|
scratch/
|
||||||
|
junifer_jobs/
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
Original Authors
|
Original Authors
|
||||||
================
|
================
|
||||||
* Federico Raimondo <f.raimondo@fz-juelich.de>
|
* 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
|
# junifer - JUelich NeuroImaging FEature extractoR
|
||||||
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
|
## Repository Organization
|
||||||
|
|
||||||
* `docs`: Documentation, built using sphinx.
|
* `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_`.
|
* `examples`: Examples, using sphinx-gallery. File names of examples that create visual output must start with `plot_`, otherwise, with `run_`.
|
||||||
* `junifer`: Main library directory
|
* `junifer`: Main library directory.
|
||||||
* `api`: User API module
|
* `api`: User API module.
|
||||||
* `data`: Module that handles data required for the library to work (e.g. atlases)
|
* `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.
|
* `datagrabber`: DataGrabber module.
|
||||||
* `datareader`: DataReader module.
|
* `datareader`: DataReader module.
|
||||||
* `markers`: Markers module.
|
* `markers`: Markers module.
|
||||||
* `pipeline`: Pipeline module.
|
* `pipeline`: Pipeline module.
|
||||||
* `preprocess`: Preprocessing module.
|
* `preprocess`: Preprocessing module.
|
||||||
* `storage`: Storage module.
|
* `storage`: Storage module.
|
||||||
|
* `testing`: Testing components module.
|
||||||
* `utils`: Utilities module (e.g. logging)
|
* `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
|
.. _Fede Raimondo: https://fraimondo.github.io
|
||||||
.. _Kaustubh Patil: https://github.com/kaurao
|
.. _Kaustubh Patil: https://github.com/kaurao
|
||||||
.. _Leonard Sasse: https://github.com/LeSasse
|
.. _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
|
- "Bugs" for bug fixes
|
||||||
- "API changes" for backward-incompatible changes
|
- "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:
|
||||||
|
|
||||||
Current (0.0.0.dev)
|
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
|
.. 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::
|
.. toctree::
|
||||||
|
:numbered:
|
||||||
:maxdepth: 2
|
:maxdepth: 2
|
||||||
:caption: Contents:
|
:caption: Contents:
|
||||||
|
|
||||||
installation
|
installation
|
||||||
api
|
understanding/index.rst
|
||||||
|
builtin
|
||||||
auto_examples/index.rst
|
auto_examples/index.rst
|
||||||
|
api/index.rst
|
||||||
|
contribution
|
||||||
maintaining
|
maintaining
|
||||||
|
faq
|
||||||
whats_new
|
whats_new
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -21,4 +35,3 @@ Indices and tables
|
||||||
* :ref:`genindex`
|
* :ref:`genindex`
|
||||||
* :ref:`modindex`
|
* :ref:`modindex`
|
||||||
* :ref:`search`
|
* :ref:`search`
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,91 +1,53 @@
|
||||||
.. include:: links.inc
|
.. include:: links.inc
|
||||||
|
|
||||||
Installing
|
Installing junifer
|
||||||
==========
|
==================
|
||||||
|
|
||||||
|
|
||||||
Requirements
|
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
|
Depending on the installation method, these packages might be installed automatically.
|
||||||
automatically.
|
|
||||||
|
|
||||||
Installing
|
Installation
|
||||||
^^^^^^^^^^
|
^^^^^^^^^^^^
|
||||||
There are different ways to install junifer:
|
Depending on your use-case, junifer can be installed differently:
|
||||||
|
|
||||||
* Install the :ref:`install_latest_release`. This is the most suitable approach
|
* Install the :ref:`install_latest_release`. This is the most suitable approach
|
||||||
for most end users.
|
for end users.
|
||||||
* Install the :ref:`install_latest_development`. This version will have the
|
* Install from :ref:`install_development_git`. This is the most suitable approach
|
||||||
latest features. However, it is still under development and not yet
|
for developers.
|
||||||
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.
|
|
||||||
|
|
||||||
|
|
||||||
Either way, we strongly recommend using virtual environments:
|
Either way, we strongly recommend using `virtual environments <https://realpython.com/python-virtual-environments-a-primer>`_.
|
||||||
|
|
||||||
* `venv`_
|
|
||||||
* `conda env`_
|
|
||||||
|
|
||||||
|
|
||||||
.. _install_latest_release:
|
.. _install_latest_release:
|
||||||
|
|
||||||
Latest release
|
Stable release
|
||||||
--------------
|
--------------
|
||||||
|
|
||||||
We have packaged junifer and published it in PyPi, so you can just install it
|
Use ``pip`` to install julearn from `PyPI <https://pypi.org>`_, like so:
|
||||||
with `pip`.
|
|
||||||
|
|
||||||
.. code-block:: bash
|
.. code-block:: bash
|
||||||
|
|
||||||
pip install -U junifer
|
pip install 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
|
|
||||||
|
|
||||||
|
|
||||||
.. _install_development_git:
|
.. _install_development_git:
|
||||||
|
|
||||||
Local git repository (for developers)
|
Local Git repository
|
||||||
-------------------------------------
|
--------------------
|
||||||
First, make sure that you have all the dependencies installed:
|
|
||||||
|
|
||||||
Then, clone `junifer Github`_ repository in a folder of your choice:
|
Follow the `detailed contribution guidelines <contribution.rst>`_.
|
||||||
|
|
||||||
.. 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``.
|
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@
|
||||||
|
|
||||||
.. _`AML`: https://www.fz-juelich.de/inm/inm-7/EN/Forschung/Applied%20Machine%20Learning/_node.html
|
.. _`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
|
.. _`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`: https://pandas.pydata.org
|
||||||
.. _`pandas.DataFrame` : https://pandas.pydata.org/pandas-docs/stable/reference/api/pandas.DataFrame.html
|
.. _`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/
|
.. _`setuptools_scm`: https://github.com/pypa/setuptools_scm/
|
||||||
|
|
||||||
.. _`sphinx gallery`: https://sphinx-gallery.github.io/stable/index.html
|
.. _`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:
|
else:
|
||||||
slug = 'issues/' + text
|
slug = 'issues/' + text
|
||||||
text = '#' + text
|
text = '#' + text
|
||||||
ref = 'https://github.com/juaml/julearn/' + slug
|
ref = 'https://github.com/juaml/junifer/' + slug
|
||||||
set_classes(options)
|
set_classes(options)
|
||||||
node = reference(rawtext, text, refuri=ref, **options)
|
node = reference(rawtext, text, refuri=ref, **options)
|
||||||
return [node], []
|
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
|
HCP FC Extraction
|
||||||
======================
|
======================
|
||||||
|
|
||||||
Authors: Leonard Sasse
|
Authors: Leonard Sasse
|
||||||
License: BSD 3 clause
|
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 = {
|
custom_confound_strategy = {
|
||||||
'filter': 'butterworth',
|
"filter": "butterworth",
|
||||||
'detrend': True,
|
"detrend": True,
|
||||||
'high_pass': 0.01,
|
"high_pass": 0.01,
|
||||||
'low_pass': 0.08,
|
"low_pass": 0.08,
|
||||||
'standardize': True,
|
"standardize": True,
|
||||||
'confounds': ['csf', 'wm', 'gsr'],
|
"confounds": ["csf", "wm", "gsr"],
|
||||||
'derivatives': True,
|
"derivatives": True,
|
||||||
'squares': True,
|
"squares": True,
|
||||||
'other': []
|
"other": [],
|
||||||
}
|
}
|
||||||
|
|
||||||
markers = [
|
markers = [
|
||||||
{'name': 'Power264_FCPearson',
|
{
|
||||||
'kind': 'FunctionalConnectivity',
|
"name": "Power264_FCPearson",
|
||||||
'atlas': 'Power264',
|
"kind": "FunctionalConnectivity",
|
||||||
'method': 'Pearson',
|
"atlas": "Power264",
|
||||||
'confound_strategy': 'Params36'},
|
"method": "Pearson",
|
||||||
{'name': 'Schaefer400x17_FCPearson',
|
"confound_strategy": "Params36",
|
||||||
'kind': 'FunctionalConnectivity',
|
},
|
||||||
'atlas': 'Schaefer400x17',
|
{
|
||||||
'method': 'Pearson',
|
"name": "Schaefer400x17_FCPearson",
|
||||||
'confound_strategy': 'Params24'},
|
"kind": "FunctionalConnectivity",
|
||||||
{'name': 'Power264_FCSpearman',
|
"atlas": "Schaefer400x17",
|
||||||
'kind': 'FunctionalConnectivity',
|
"method": "Pearson",
|
||||||
'atlas': 'Power264',
|
"confound_strategy": "Params24",
|
||||||
'method': 'Spearman',
|
},
|
||||||
'confound_strategy': 'ICAAROMA'},
|
{
|
||||||
{'name': 'Schaefer400x17_FCSpearman',
|
"name": "Power264_FCSpearman",
|
||||||
'kind': 'FunctionalConnectivity',
|
"kind": "FunctionalConnectivity",
|
||||||
'atlas': 'Schaefer400x17',
|
"atlas": "Power264",
|
||||||
'method': 'Spearman',
|
"method": "Spearman",
|
||||||
'confound_strategy': 'path/to/predefined/confound_file.tsv'},
|
"confound_strategy": "ICAAROMA",
|
||||||
{'name': 'Schaefer400x17_FCSpearman',
|
},
|
||||||
'kind': 'FunctionalConnectivity',
|
{
|
||||||
'atlas': 'Schaefer400x17',
|
"name": "Schaefer400x17_FCSpearman",
|
||||||
'method': 'Spearman',
|
"kind": "FunctionalConnectivity",
|
||||||
'confound_strategy': custom_confound_strategy}
|
"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 = {
|
storage = {
|
||||||
'modality': 'fMRI',
|
"kind": "SQLiteFeatureStorage",
|
||||||
'preprocessed': 'ICA+FIX',
|
"uri": "/data/project/juniferexample",
|
||||||
'space': 'volumetric',
|
|
||||||
}
|
}
|
||||||
|
|
||||||
run_pipeline(
|
run(
|
||||||
workdir='/tmp',
|
workdir="/tmp",
|
||||||
datagrabber='HCPOpenAccess',
|
datagrabber=datagrabber,
|
||||||
datagrabber_params=dg_params,
|
elements=[("100408", "REST1", "LR")],
|
||||||
element=('100408', 'REST1', "LR"),
|
|
||||||
markers=markers,
|
markers=markers,
|
||||||
storage='SQLDataFrameStorage',
|
storage=storage,
|
||||||
storage_params={'outpath': '/data/project/juniferexample'},
|
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -8,28 +8,36 @@ License: BSD 3 clause
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
from junifer.api import run_pipeline
|
from junifer.api import run
|
||||||
|
|
||||||
|
|
||||||
markers = [
|
markers = [
|
||||||
{'name': 'Schaefer1000x7_TrimMean80',
|
{
|
||||||
'kind': 'ParcelAggregation',
|
"name": "Schaefer1000x7_TrimMean80",
|
||||||
'atlas': 'Schaefer1000x7',
|
"kind": "ParcelAggregation",
|
||||||
'method': 'trimmean80'},
|
"atlas": "Schaefer1000x7",
|
||||||
{'name': 'Schaefer1000x7_Mean',
|
"method": "trim_mean",
|
||||||
'kind': 'ParcelAggregation',
|
"method_params": {"proportiontocut": 0.2},
|
||||||
'atlas': 'Schaefer1000x7',
|
},
|
||||||
'method': 'mean'},
|
{
|
||||||
{'name': 'Schaefer1000x7_Std',
|
"name": "Schaefer1000x7_Mean",
|
||||||
'kind': 'ParcelAggregation',
|
"kind": "ParcelAggregation",
|
||||||
'atlas': 'Schaefer1000x7',
|
"atlas": "Schaefer1000x7",
|
||||||
'method': 'std'}
|
"method": "mean",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "Schaefer1000x7_Std",
|
||||||
|
"kind": "ParcelAggregation",
|
||||||
|
"atlas": "Schaefer1000x7",
|
||||||
|
"method": "std",
|
||||||
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
run_pipeline(
|
run(
|
||||||
workdir='/tmp',
|
workdir="/tmp",
|
||||||
datagrabber='JuselessUKBVBM',
|
datagrabber="JuselessUKBVBM",
|
||||||
element=('sub-1627474', 'ses-2'),
|
elements=("sub-1627474", "ses-2"),
|
||||||
markers=markers,
|
markers=markers,
|
||||||
storage='SQLDataFrameStorage',
|
storage="SQLDataFrameStorage",
|
||||||
storage_params={'outpath': '/data/project/juniferexample'},
|
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
|
License: BSD 3 clause
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from junifer.datagrabber.base import BIDSDataladDataGrabber
|
from junifer.datagrabber import PatternDataladDataGrabber
|
||||||
from junifer.utils import configure_logging
|
from junifer.utils import configure_logging
|
||||||
|
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Set the logging level to info to see extra information
|
# 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,
|
# The BIDS datagrabber requires three parameters: the types of data we want,
|
||||||
# and the specific pattern that matches each type.
|
# the specific pattern that matches each type, and the variables that will be
|
||||||
types = ['T1w', 'bold']
|
# replaced int he patterns.
|
||||||
|
types = ["T1w", "bold"]
|
||||||
patterns = {
|
patterns = {
|
||||||
'T1w': 'anat/{subject}_T1w.nii.gz',
|
"T1w": "{subject}/anat/{subject}_T1w.nii.gz",
|
||||||
'bold': 'func/{subject}_task-rest_bold.nii.gz'
|
"bold": "{subject}/func/{subject}_task-rest_bold.nii.gz",
|
||||||
}
|
}
|
||||||
|
replacements = ["subject"]
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Additionally, a datalad datagrabber requires the URI of the remote sibling
|
# Additionally, a datalad datagrabber requires the URI of the remote sibling
|
||||||
# and the location of the dataset within the remote sibling.
|
# and the location of the dataset within the remote sibling.
|
||||||
repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids'
|
repo_uri = "https://gin.g-node.org/juaml/datalad-example-bids"
|
||||||
rootdir = 'example_bids'
|
rootdir = "example_bids"
|
||||||
|
|
||||||
###############################################################################
|
###############################################################################
|
||||||
# Now we can use the datagrabber within a `with` context
|
# Now we can use the datagrabber within a `with` context
|
||||||
# One thing we can do with any datagrabber is iterate over the elements.
|
# One thing we can do with any datagrabber is iterate over the elements.
|
||||||
# In this case, each element of the datagrabber is one session.
|
# In this case, each element of the datagrabber is one session.
|
||||||
with BIDSDataladDataGrabber(rootdir=rootdir, types=types,
|
with PatternDataladDataGrabber(
|
||||||
patterns=patterns, uri=repo_uri) as dg:
|
rootdir=rootdir,
|
||||||
|
types=types,
|
||||||
|
patterns=patterns,
|
||||||
|
uri=repo_uri,
|
||||||
|
replacements=replacements,
|
||||||
|
) as dg:
|
||||||
for elem in dg:
|
for elem in dg:
|
||||||
print(elem)
|
print(elem)
|
||||||
|
|
||||||
|
|
@ -46,7 +53,12 @@ with BIDSDataladDataGrabber(rootdir=rootdir, types=types,
|
||||||
# Another feature of the datagrabber is the ability to get a specific
|
# 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
|
# 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).
|
# paths for the two types of data we want (T1w and bold).
|
||||||
with BIDSDataladDataGrabber(rootdir=rootdir, types=types,
|
with PatternDataladDataGrabber(
|
||||||
patterns=patterns, uri=repo_uri) as dg:
|
rootdir=rootdir,
|
||||||
sub01 = dg['sub-01']
|
types=types,
|
||||||
|
patterns=patterns,
|
||||||
|
uri=repo_uri,
|
||||||
|
replacements=replacements,
|
||||||
|
) as dg:
|
||||||
|
sub01 = dg["sub-01"]
|
||||||
print(sub01)
|
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 api
|
||||||
from . import utils
|
from . import configs
|
||||||
|
from . import data
|
||||||
from . import datagrabber
|
from . import datagrabber
|
||||||
|
from . import datareader
|
||||||
from . import markers
|
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>
|
# 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>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from . pipeline import register
|
|
||||||
|
from .registry import register
|
||||||
|
|
||||||
|
|
||||||
def register_datagrabber(klass):
|
def register_datagrabber(klass: type) -> type:
|
||||||
"""Datagrabber decorator.
|
"""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
|
Parameters
|
||||||
----------
|
----------
|
||||||
|
|
@ -17,7 +21,60 @@ def register_datagrabber(klass):
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
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
|
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 ..api.decorators import register_datagrabber
|
||||||
|
from ..datagrabber import PatternDataladDataGrabber
|
||||||
|
|
||||||
|
|
||||||
@register_datagrabber
|
@register_datagrabber
|
||||||
class JuselessUKBVBM(DataladDataGrabber):
|
class JuselessDataladUKBVBM(PatternDataladDataGrabber):
|
||||||
"""Juseless UKB VMG DataGrabber class.
|
"""Juseless UKB VBM DataGrabber class.
|
||||||
|
|
||||||
Implements a DataGrabber to access the UKB VBM data in Juseless.
|
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):
|
def __init__(self, datadir: Union[str, Path, None] = None) -> None:
|
||||||
"""Initialize a JuselessUKBVBM object.
|
"""Initialize the class."""
|
||||||
|
uri = "ria+http://ukb.ds.inm7.de#~cat_m0wp1"
|
||||||
Parameters
|
rootdir = "m0wp1"
|
||||||
----------
|
types = ["VBM_GM"]
|
||||||
datadir : str or Path
|
replacements = ["subject", "session"]
|
||||||
That directory where the datalad dataset will be cloned. If None,
|
patterns = {"VBM_GM": "m0wp1sub-{subject}_ses-{session}_T1w.nii.gz"}
|
||||||
(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']
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
types=types, datadir=datadir, uri=uri, rootdir=rootdir)
|
types=types,
|
||||||
|
datadir=datadir,
|
||||||
def get_elements(self):
|
uri=uri,
|
||||||
"""Get the list of subjects in the dataset.
|
rootdir=rootdir,
|
||||||
|
replacements=replacements,
|
||||||
Returns
|
patterns=patterns,
|
||||||
-------
|
)
|
||||||
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
|
|
||||||
|
|
|
||||||
|
|
@ -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 socket
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.configs.juseless import JuselessUKBVBM
|
from junifer.configs.juseless import JuselessDataladUKBVBM
|
||||||
|
from junifer.datagrabber.hcp import DataladHCP1200
|
||||||
if socket.gethostname() != 'juseless':
|
from junifer.utils.logging import configure_logging
|
||||||
pytest.skip('This tests are only for juseless', allow_module_level=True)
|
|
||||||
|
|
||||||
|
|
||||||
def test_juselessukbvbm_datagrabber():
|
# Check if the test is running on juseless
|
||||||
with JuselessUKBVBM() as dg:
|
if socket.gethostname() != "juseless":
|
||||||
out = dg[('sub-2670511', 'ses-2')]
|
pytest.skip("These tests are only for juseless", allow_module_level=True)
|
||||||
assert 'VBM_GM' in out
|
|
||||||
assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz'
|
configure_logging(level="DEBUG")
|
||||||
assert out['VBM_GM'].exists()
|
|
||||||
|
|
||||||
|
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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from pathlib import Path
|
|
||||||
import io
|
import io
|
||||||
import requests
|
import shutil
|
||||||
import numpy as np
|
import tempfile
|
||||||
import pandas as pd
|
import zipfile
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import requests
|
||||||
from nilearn import datasets
|
from nilearn import datasets
|
||||||
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nibabel import Nifti1Image
|
||||||
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
A dictionary containing all supported atlases and their respective valid
|
A dictionary containing all supported atlases and their respective valid
|
||||||
parameters.
|
parameters.
|
||||||
|
|
@ -23,150 +36,198 @@ Optional keys:
|
||||||
* 'valid_resolutions': a list of valid resolutions for the atlas (e.g. [1, 2])
|
* 'valid_resolutions': a list of valid resolutions for the atlas (e.g. [1, 2])
|
||||||
|
|
||||||
"""
|
"""
|
||||||
_available_atlases = {
|
# TODO: have separate dictionary for built-in
|
||||||
'SUITxSUIT': {
|
_available_atlases: Dict[str, Dict[Any, Any]] = {
|
||||||
'family': 'SUIT',
|
"SUITxSUIT": {"family": "SUIT", "space": "SUIT"},
|
||||||
'sace': 'SUIT'
|
"SUITxMNI": {"family": "SUIT", "space": "MNI"},
|
||||||
},
|
|
||||||
'SUITxMNI': {
|
|
||||||
'family': 'SUIT',
|
|
||||||
'space': 'MNI'
|
|
||||||
},
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Add Schaefer atlas info
|
||||||
for n_rois in range(100, 1001, 100):
|
for n_rois in range(100, 1001, 100):
|
||||||
for t_net in [7, 17]:
|
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] = {
|
_available_atlases[t_name] = {
|
||||||
'family': 'Schaefer',
|
"family": "Schaefer",
|
||||||
'n_rois': n_rois,
|
"n_rois": n_rois,
|
||||||
'yeo_networks': t_net,
|
"yeo_networks": t_net,
|
||||||
'valid_resolutions': [1, 2]
|
|
||||||
}
|
}
|
||||||
|
# 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.
|
"""Register a custom user atlas.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
name : str
|
name : str
|
||||||
The name of the atlas.
|
The name of the atlas.
|
||||||
atlas_path : str
|
atlas_path : str or pathlib.Path
|
||||||
The path to the atlas file.
|
The path to the atlas file.
|
||||||
atl_labels : list(str)
|
atl_labels : list of str
|
||||||
The list of labels for the atlas.
|
The list of labels for the atlas.
|
||||||
overwrite : bool
|
overwrite : bool, optional
|
||||||
If True, overwrite an existing atlas with the same name. Defaults to
|
If True, overwrite an existing atlas with the same name.
|
||||||
False.
|
Does not apply to built-in atlases (default False).
|
||||||
|
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
ValueError
|
ValueError
|
||||||
If the atlas name is already registered and overwrite is set to False
|
If the atlas name is already registered and overwrite is set to False
|
||||||
or if the atlas name is a built-in atlas.
|
or if the atlas name is a built-in atlas.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
# Check for attempt of overwriting built-in atlases
|
||||||
if name in _available_atlases:
|
if name in _available_atlases:
|
||||||
if overwrite is True:
|
if overwrite is True:
|
||||||
logger.info(f'Overwritting {name} atlas')
|
logger.info(f"Overwriting {name} atlas")
|
||||||
if _available_atlases[name]['family'] != 'CustomUserAtlas':
|
if _available_atlases[name]["family"] != "CustomUserAtlas":
|
||||||
raise_error(
|
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:
|
else:
|
||||||
raise_error(
|
raise_error(
|
||||||
f'Atlas {name} already registered. Set `overwrite=True` to '
|
f"Atlas {name} already registered. Set `overwrite=True` to "
|
||||||
'update its value.')
|
"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] = {
|
_available_atlases[name] = {
|
||||||
'path': atlas_path, 'labels': atl_labels,
|
"path": str(atlas_path.absolute()),
|
||||||
'family': 'CustomUserAtlas'}
|
"labels": atl_labels,
|
||||||
|
"family": "CustomUserAtlas",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def list_atlases():
|
def list_atlases() -> List[str]:
|
||||||
"""
|
"""List all the available atlases.
|
||||||
List all the available atlases.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
out : list(str) or dict
|
list of str
|
||||||
A list or dict with all available atlases.
|
A list with all available atlases.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
return sorted(_available_atlases.keys())
|
return sorted(_available_atlases.keys())
|
||||||
|
|
||||||
|
|
||||||
def _check_resolution(resolution, valid_resolution):
|
# def _check_resolution(resolution, valid_resolution):
|
||||||
if resolution is None:
|
# if resolution is None:
|
||||||
return None
|
# return None
|
||||||
if resolution not in valid_resolution:
|
# if resolution not in valid_resolution:
|
||||||
raise ValueError(f'Invalid resolution: {resolution}')
|
# raise ValueError(f'Invalid resolution: {resolution}')
|
||||||
return resolution
|
# return resolution
|
||||||
|
|
||||||
|
|
||||||
def load_atlas(name, atlas_dir=None, resolution=None, path_only=False,
|
# TODO: keyword arguments are not passed, check
|
||||||
**kwargs):
|
def load_atlas(
|
||||||
"""
|
name: str,
|
||||||
Loads a brain atlas (including a label file).
|
atlas_dir: Union[str, Path, None] = None,
|
||||||
If it is built-in atlas and file is not present in the `atlas_dir`
|
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.
|
directory, it will be downloaded.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
name : str
|
name : str
|
||||||
The name of the atlas.
|
The name of the atlas. Check valid options by calling `list_atlases`.
|
||||||
Check valid options by calling `list_atlases`.
|
atlas_dir : str or pathlib.Path, optional
|
||||||
atlas_dir: path
|
Path where the atlas files are stored. The default location is
|
||||||
Path where the atlas files are stored.
|
"$HOME/junifer/data/atlas" (default None).
|
||||||
Defaults to: $HOME/junifer/data/atlas
|
resolution : float, optional
|
||||||
resolution : int
|
The desired resolution of the atlas to load. If it is not available,
|
||||||
The (desired) resolution of the atlas to load. If its not available,
|
|
||||||
the closest resolution will be loaded. Preferably, use a resolution
|
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
|
||||||
path_only : bool
|
(default None).
|
||||||
If True, the atlas image will not be loaded.
|
path_only : bool, optional
|
||||||
|
If True, the atlas image will not be loaded (default False).
|
||||||
|
|
||||||
Parameters (optional, atlas dependent)
|
Extra Parameters
|
||||||
--------------------------------------
|
----------------
|
||||||
Use to specify atlas specific keyword arguments. .
|
Use to specify atlas specific keyword arguments.
|
||||||
|
|
||||||
Schaefer :
|
- Schaefer :
|
||||||
n_rois (required) : int
|
n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000}
|
||||||
Granularity of atlas to be used. Valid values: between 100 and 1000
|
Granularity of atlas to be used.
|
||||||
(included) in steps of 100.
|
yeo_network : {7, 17}, optional
|
||||||
yeo_network (optional) : int
|
Number of yeo networks to use (default 7).
|
||||||
Number of yeo networks to use. Valid values: 7, 17. Defaults to 7.
|
- Tian :
|
||||||
|
scale : {1, 2, 3, 4}
|
||||||
Tian :
|
Scale of atlas (defines granularity).
|
||||||
# TODO add
|
space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional
|
||||||
|
Space of atlas (default "MNI6thgeneration"). (For more information
|
||||||
SUIT :
|
see https://github.com/yetianmed/subcortex)
|
||||||
space (optional) : str
|
magneticfield : {"3T", "7T"}, optional
|
||||||
Space of atlas can be either 'MNI' or 'SUIT' (for more information
|
Magnetic field (default "3T").
|
||||||
see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to
|
- SUIT :
|
||||||
'MNI'.
|
space : {"MNI", "SUIT"}, optional
|
||||||
|
Space of atlas (default "MNI"). (For more information
|
||||||
|
see http://www.diedrichsenlab.org/imaging/suit.htm).
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
atlas_img : niimg-like object or None
|
niimg-like object or None
|
||||||
Loaded atlas image.
|
Loaded atlas image.
|
||||||
atlas_labels : List of str
|
list of str
|
||||||
Atlas labels.
|
Atlas labels.
|
||||||
atlas_fname : Path
|
pathlib.Path
|
||||||
File path to the atlas image.
|
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]
|
atlas_definition = _available_atlases[name].copy()
|
||||||
t_family = atlas_definition.pop('family')
|
t_family = atlas_definition.pop("family")
|
||||||
|
|
||||||
if t_family == 'CustomUserAtlas':
|
if t_family == "CustomUserAtlas":
|
||||||
atlas_fname = atlas_definition['path']
|
atlas_fname = Path(atlas_definition["path"])
|
||||||
atlas_labels = atlas_definition['labels']
|
atlas_labels = atlas_definition["labels"]
|
||||||
else:
|
else:
|
||||||
# retrieve atlases by passing arguments on to _retrieve_atlas()
|
# retrieve atlases by passing arguments on to _retrieve_atlas()
|
||||||
atlas_fname, atlas_labels = _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(
|
logger.info(f"Loading atlas {str(atlas_fname.absolute())}")
|
||||||
f'Loading atlas {atlas_fname.as_posix()}') # type: ignore
|
|
||||||
|
|
||||||
atlas_img = None
|
atlas_img = None
|
||||||
if path_only is False:
|
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
|
return atlas_img, atlas_labels, atlas_fname
|
||||||
|
|
||||||
|
|
||||||
def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs):
|
def _retrieve_atlas(
|
||||||
"""
|
family: str,
|
||||||
Retrieves a brain atlas object either from nilearn or a specified online
|
atlas_dir: Union[str, Path, None] = None,
|
||||||
source. Only returns one atlas per call. Call function multiple times for
|
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
|
different parameter specifications. Only retrieves atlas if it is not yet
|
||||||
in atlas_dir.
|
in atlas_dir.
|
||||||
|
|
||||||
Parameters (required)
|
Parameters
|
||||||
---------------------
|
----------
|
||||||
family : str
|
family : str
|
||||||
Specify by name of atlas family, e.g. 'Schaefer'.
|
The name of the atlas family, e.g. 'Schaefer'.
|
||||||
atlas_dir: path
|
atlas_dir : str or pathlib.Path, optional
|
||||||
Path to where to store the retrieved atlas file.
|
Path where the retrieved atlas file is stored. The default location is
|
||||||
Defaults to: $HOME/junifer/data/atlas
|
"$HOME/junifer/data/atlas" (default None).
|
||||||
resolution : int
|
resolution : float, optional
|
||||||
The (desired) resolution of the atlas to load. If its not available,
|
The desired resolution of the atlas to load. If it is not available,
|
||||||
the closest resolution will be loaded. Preferably, use a resolution
|
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)
|
Extra Parameters
|
||||||
--------------------------------------
|
----------------
|
||||||
Use to specify atlas specific keyword arguments
|
**kwargs
|
||||||
|
Use to specify atlas specific keyword arguments.
|
||||||
|
|
||||||
Schaefer :
|
- Schaefer :
|
||||||
n_rois (required) : int
|
n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000}
|
||||||
Granularity of atlas to be used. Valid values: between 100 and 1000
|
Granularity of atlas to be used.
|
||||||
(included) in steps of 100.
|
yeo_network : {7, 17}, optional
|
||||||
yeo_network (optional) : int
|
Number of yeo networks to use (default 7).
|
||||||
Number of yeo networks to use [7 or 17]. Defaults to 7.
|
- Tian :
|
||||||
SUIT :
|
scale : {1, 2, 3, 4}
|
||||||
space (optional) : str
|
Scale of atlas (defines granularity).
|
||||||
Space of atlas can be either 'MNI' or 'SUIT' (for more information
|
space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional
|
||||||
see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to
|
Space of atlas (default "MNI6thgeneration"). (For more
|
||||||
'MNI'.
|
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
|
Returns
|
||||||
-------
|
-------
|
||||||
atlas_fname : Path
|
pathlib.Path
|
||||||
File path to the atlas image.
|
File path to the atlas image.
|
||||||
atlas_labels : List of str
|
list of str
|
||||||
Atlas labels.
|
Atlas labels.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If the atlas name is invalid.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
if atlas_dir is None:
|
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)
|
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.")
|
logger.info(f"Fetching one of {family} atlas.")
|
||||||
|
|
||||||
# retrieval details per atlas
|
# Retrieval details per atlas
|
||||||
if family == 'Schaefer':
|
if family == "Schaefer":
|
||||||
atlas_fname, atl_labels = \
|
atlas_fname, atl_labels = _retrieve_schaefer(
|
||||||
_retrieve_schaefer(atlas_dir, **kwargs)
|
atlas_dir=atlas_dir, resolution=resolution, **kwargs
|
||||||
elif family == 'SUIT':
|
)
|
||||||
atlas_fname, atl_labels = \
|
elif family == "SUIT":
|
||||||
_retrieve_suit(atlas_dir, **kwargs)
|
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:
|
else:
|
||||||
raise_error(
|
raise_error(f"The provided atlas name {family} cannot be retrieved.")
|
||||||
f"The provided atlas name {family} cannot be retrieved. ")
|
|
||||||
|
|
||||||
return atlas_fname, atl_labels
|
return atlas_fname, atl_labels
|
||||||
|
|
||||||
|
|
||||||
def _closest_resolution(resolution, valid_resolution):
|
def _closest_resolution(
|
||||||
closest = None
|
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):
|
if not isinstance(valid_resolution, np.ndarray):
|
||||||
valid_resolution = np.array(valid_resolution)
|
valid_resolution = np.array(valid_resolution)
|
||||||
|
|
||||||
if resolution is None:
|
if resolution is None:
|
||||||
logger.info('Resolution set to None, using highest resolution.')
|
logger.info("Resolution set to None, using highest resolution.")
|
||||||
closest = np.min(valid_resolution)
|
closest = np.min(valid_resolution)
|
||||||
elif any(x <= resolution for x in valid_resolution):
|
elif any(x <= resolution for x in valid_resolution):
|
||||||
# Case 1: get the highest closest resolution
|
# Case 1: get the highest closest resolution
|
||||||
closest = np.max(valid_resolution[valid_resolution <= resolution])
|
closest = np.max(valid_resolution[valid_resolution <= resolution])
|
||||||
|
|
@ -254,11 +362,46 @@ def _closest_resolution(resolution, valid_resolution):
|
||||||
return closest
|
return closest
|
||||||
|
|
||||||
|
|
||||||
def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7):
|
def _retrieve_schaefer(
|
||||||
logger.info('Atlas parameters:')
|
atlas_dir: Path,
|
||||||
logger.info(f'\tn_rois: {n_rois}')
|
resolution: Optional[float] = None,
|
||||||
logger.info(f'\tyeo_network: {yeo_network}')
|
n_rois: Optional[int] = None,
|
||||||
logger.info(f'\tresolution: {resolution}')
|
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_n_rois = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000]
|
||||||
_valid_networks = [7, 17]
|
_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:
|
if n_rois not in _valid_n_rois:
|
||||||
raise_error(
|
raise_error(
|
||||||
f'The parameter `n_rois` ({n_rois}) needs to be one of the '
|
f"The parameter `n_rois` ({n_rois}) needs to be one of the "
|
||||||
f'following: {_valid_n_rois}')
|
f"following: {_valid_n_rois}"
|
||||||
if yeo_network not in _valid_networks:
|
)
|
||||||
|
if yeo_networks not in _valid_networks:
|
||||||
raise_error(
|
raise_error(
|
||||||
f'The parameter `yeo_network` ({yeo_network}) needs to be one of '
|
f"The parameter `yeo_networks` ({yeo_networks}) needs to be one "
|
||||||
f'the following: {_valid_networks}')
|
f"of the following: {_valid_networks}"
|
||||||
|
)
|
||||||
|
|
||||||
resolution = _closest_resolution(resolution, _valid_resolutions)
|
resolution = _closest_resolution(resolution, _valid_resolutions)
|
||||||
|
|
||||||
# define file names
|
# define file names
|
||||||
atlas_fname = atlas_dir / 'schaefer_2018' / (
|
atlas_fname = (
|
||||||
f'Schaefer2018_{n_rois}Parcels_{yeo_network}Networks_order_'
|
atlas_dir
|
||||||
f'FSLMNI152_{resolution}mm.nii.gz')
|
/ "schaefer_2018"
|
||||||
atlas_lname = atlas_dir / 'schaefer_2018' / (
|
/ (
|
||||||
f'Schaefer2018_{n_rois}Parcels_{yeo_network}Networks_order.txt')
|
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()):
|
if not (atlas_fname.exists() and atlas_lname.exists()):
|
||||||
logger.info(
|
logger.info(
|
||||||
'At least one of the atlas files is missing. '
|
"At least one of the atlas files is missing. "
|
||||||
'Fetching using nilearn.')
|
"Fetching using nilearn."
|
||||||
|
)
|
||||||
datasets.fetch_atlas_schaefer_2018(
|
datasets.fetch_atlas_schaefer_2018(
|
||||||
n_rois=n_rois,
|
n_rois=n_rois,
|
||||||
yeo_networks=yeo_network,
|
yeo_networks=yeo_networks,
|
||||||
resolution_mm=resolution,
|
resolution_mm=resolution, # type: ignore we know it's 1 or 2
|
||||||
data_dir=atlas_dir.as_posix())
|
data_dir=str(atlas_dir.absolute()),
|
||||||
|
)
|
||||||
|
|
||||||
if not (atlas_fname.exists() and atlas_lname.exists()):
|
if not (
|
||||||
raise_error('There was a problem fetching the atlases.')
|
atlas_fname.exists() and atlas_lname.exists()
|
||||||
|
): # pragma: no cover
|
||||||
|
raise_error("There was a problem fetching the atlases.")
|
||||||
|
|
||||||
# Load labels
|
# Load labels
|
||||||
labels = [
|
labels = [
|
||||||
'_'.join(x.split('_')[1:])
|
"_".join(x.split("_")[1:])
|
||||||
for x in pd.read_csv(
|
for x in pd.read_csv(atlas_lname, sep="\t", header=None)
|
||||||
atlas_lname, sep='\t', header=None).iloc[:, 1].to_list()
|
.iloc[:, 1]
|
||||||
|
.to_list()
|
||||||
]
|
]
|
||||||
|
|
||||||
return atlas_fname, labels
|
return atlas_fname, labels
|
||||||
|
|
||||||
|
|
||||||
def _retrieve_suit(out_dir, resolution, space='MNI'):
|
def _retrieve_tian(
|
||||||
logger.info('Atlas parameters:')
|
atlas_dir: Path,
|
||||||
logger.info(f'\tspace: {space}')
|
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
|
# check validity of atlas parameters
|
||||||
if space not in _valid_spaces:
|
if space not in _valid_spaces:
|
||||||
raise_error(
|
raise_error(
|
||||||
f'The parameter `space` ({space}) needs to be one of the '
|
f"The parameter `space` ({space}) needs to be one of the "
|
||||||
f'following: {_valid_spaces}')
|
f"following: {_valid_spaces}"
|
||||||
|
)
|
||||||
|
|
||||||
# TODO: Validate this with Vera
|
# TODO: Validate this with Vera
|
||||||
_valid_resolutions = [1]
|
_valid_resolutions = [1]
|
||||||
|
|
@ -324,45 +675,52 @@ def _retrieve_suit(out_dir, resolution, space='MNI'):
|
||||||
resolution = _closest_resolution(resolution, _valid_resolutions)
|
resolution = _closest_resolution(resolution, _valid_resolutions)
|
||||||
|
|
||||||
# define file names
|
# define file names
|
||||||
atlas_fname = out_dir / 'SUIT' / (
|
atlas_fname = (
|
||||||
f'SUIT_{space}Space_{resolution}mm.nii')
|
atlas_dir / "SUIT" / (f"SUIT_{space}Space_{resolution}mm.nii")
|
||||||
atlas_lname = out_dir / 'SUIT' / (
|
)
|
||||||
f'SUIT_{space}Space_{resolution}mm.tsv')
|
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()):
|
if not (atlas_fname.exists() and atlas_lname.exists()):
|
||||||
logger.info(
|
atlas_fname.parent.mkdir(exist_ok=True, parents=True)
|
||||||
'At least one of the atlas files is missing. '
|
logger.info("At least one of the atlas files is missing, fetching.")
|
||||||
'Fetching.')
|
|
||||||
|
|
||||||
url_basis = (
|
url_basis = (
|
||||||
'https://github.com/DiedrichsenLab/cerebellar_atlases/blob'
|
"https://github.com/DiedrichsenLab/cerebellar_atlases/raw"
|
||||||
'/master/Diedrichsen_2009/')
|
"/master/Diedrichsen_2009/"
|
||||||
url_MNI = url_basis + 'atl-Anatom_space-MNI_dseg.nii'
|
)
|
||||||
url_SUIT = url_basis + 'atl-Anatom_space-SUIT_dseg.nii'
|
url_MNI = url_basis + "atl-Anatom_space-MNI_dseg.nii"
|
||||||
url_labels = url_basis + 'atl-Anatom.tsv'
|
url_SUIT = url_basis + "atl-Anatom_space-SUIT_dseg.nii"
|
||||||
|
url_labels = url_basis + "atl-Anatom.tsv"
|
||||||
|
|
||||||
if space == 'MNI':
|
if space == "MNI":
|
||||||
logger.info(f'Downloading {url_MNI}')
|
logger.info(f"Downloading {url_MNI}")
|
||||||
atlas_download = requests.get(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)
|
f.write(atlas_download.content)
|
||||||
elif space == 'SUIT':
|
else: # if not MNI, then SUIT
|
||||||
logger.info(f'Downloading {url_SUIT}')
|
logger.info(f"Downloading {url_SUIT}")
|
||||||
atlas_download = requests.get(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)
|
f.write(atlas_download.content)
|
||||||
|
|
||||||
labels_download = requests.get(url_labels)
|
labels_download = requests.get(url_labels)
|
||||||
labels = pd.read_csv(
|
labels = pd.read_csv(
|
||||||
io.StringIO(labels_download.content.decode("utf-8")),
|
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)
|
labels.to_csv(atlas_lname, sep="\t", index=False)
|
||||||
if not (atlas_fname.exists() and atlas_lname.exists()):
|
if (
|
||||||
raise_error('There was a problem fetching the atlases.')
|
not atlas_fname.exists() and atlas_lname.exists()
|
||||||
|
): # pragma: no cover
|
||||||
|
raise_error("There was a problem fetching the atlases.")
|
||||||
|
|
||||||
labels = pd.read_csv(
|
labels = pd.read_csv(atlas_lname, sep="\t", usecols=["name"])[
|
||||||
atlas_lname, sep='\t', usecols=['name'])['name'].to_list()
|
"name"
|
||||||
|
].to_list()
|
||||||
|
|
||||||
return atlas_fname, labels
|
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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Leonard Sasse <l.sasse@fz-juelich.de>
|
# Leonard Sasse <l.sasse@fz-juelich.de>
|
||||||
# License: AGPL
|
# 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>
|
# 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>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from pathlib import Path
|
|
||||||
import tempfile
|
|
||||||
|
|
||||||
import datalad.api as dl
|
|
||||||
from abc import ABC, abstractmethod
|
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 import logger, raise_error
|
||||||
from ..utils.logging import logger, raise_error
|
from .utils import validate_types
|
||||||
|
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
|
|
||||||
class BaseDataGrabber(ABC):
|
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
|
Attributes
|
||||||
----------
|
----------
|
||||||
datadir
|
datadir : pathlib.Path
|
||||||
types : list
|
The directory where the data is / will be stored.
|
||||||
List of data types to be grabbed.
|
|
||||||
|
|
||||||
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
|
def __init__(self, types: List[str], datadir: Union[str, Path]) -> None:
|
||||||
----------
|
"""Initialize the class."""
|
||||||
types : list of str
|
# Validate types
|
||||||
The types of data to be grabbed.
|
validate_types(types)
|
||||||
datadir : str or Path
|
# Convert str to Path
|
||||||
That directory where the data is/will be stored.
|
|
||||||
"""
|
|
||||||
_validate_types(types)
|
|
||||||
if not isinstance(datadir, Path):
|
if not isinstance(datadir, Path):
|
||||||
datadir = Path(datadir)
|
datadir = Path(datadir)
|
||||||
|
logger.debug("Initializing BaseDataGrabber")
|
||||||
|
logger.debug(f"\t_datadir = {datadir}")
|
||||||
|
logger.debug(f"\ttypes = {types}")
|
||||||
self._datadir = datadir
|
self._datadir = datadir
|
||||||
self.types = types
|
self.types = types
|
||||||
|
|
||||||
@property
|
def __iter__(self) -> Iterator:
|
||||||
def datadir(self):
|
"""Enable iterable support.
|
||||||
"""
|
|
||||||
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.
|
|
||||||
|
|
||||||
Yields
|
Yields
|
||||||
------
|
------
|
||||||
element : object
|
object
|
||||||
An element that can be indexed by the datagrabber.
|
An element that can be indexed by the datagrabber.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
for elem in self.get_elements():
|
for elem in self.get_elements():
|
||||||
yield elem
|
yield elem
|
||||||
|
|
||||||
@abstractmethod
|
# TODO: element does nothing, check
|
||||||
def __getitem__(self, element):
|
def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]:
|
||||||
raise NotImplementedError('__getitem__ not implemented')
|
"""Enable indexing support.
|
||||||
|
|
||||||
@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.
|
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
types : list of str
|
element : str or tuple
|
||||||
The types of data to be grabbed.
|
The element to be indexed. If one string is provided, it is
|
||||||
patterns : dict[str -> str]
|
assumed to be a tuple with only one item. If a tuple is provided,
|
||||||
Patterns for each type of data. The keys are the types and the
|
each item in the tuple is the value for the replacement string
|
||||||
values are the patterns. Each occurrence of the string `{subject}`
|
specified in "replacements".
|
||||||
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
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
elems : list[str]
|
dict
|
||||||
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]
|
|
||||||
Dictionary of paths for each type of data required for the
|
Dictionary of paths for each type of data required for the
|
||||||
specified element.
|
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
|
return out
|
||||||
|
|
||||||
|
def __enter__(self) -> "BaseDataGrabber":
|
||||||
@register_datagrabber
|
"""Context entry."""
|
||||||
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()
|
|
||||||
return self
|
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
|
@property
|
||||||
def datadir(self):
|
def datadir(self) -> Path:
|
||||||
return super().datadir / self._rootdir
|
"""Get data directory path.
|
||||||
|
|
||||||
def install(self):
|
Returns
|
||||||
"""Install the datalad dataset into the datadir."""
|
-------
|
||||||
self.dataset = dl.install( # type: ignore
|
pathlib.Path
|
||||||
self._datadir, source=self.uri)
|
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):
|
@abstractmethod
|
||||||
"""Remove the datalad dataset from the datadir."""
|
def get_elements(self) -> List:
|
||||||
self.dataset.remove(recursive=True)
|
"""Get elements.
|
||||||
|
|
||||||
def _dataset_get(self, out):
|
Returns
|
||||||
for _, v in out.items():
|
-------
|
||||||
self.dataset.get(v['path'])
|
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
|
raise_error(
|
||||||
the paths from the parent class and then `datalad get` each of the
|
msg="Concrete classes need to implement get_elements().",
|
||||||
files."""
|
klass=NotImplementedError,
|
||||||
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
|
|
||||||
|
|
|
||||||
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>
|
# 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>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
import pytest
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from junifer.datagrabber.base import BaseDataGrabber
|
||||||
|
|
||||||
|
|
||||||
def test_BIDSDataGrabber():
|
def test_BaseDataGrabber_abstractness() -> None:
|
||||||
"""Test BIDSDataGrabber"""
|
"""Test BaseDataGrabber is abstract base class."""
|
||||||
with pytest.raises(TypeError, match=r"types must be a list"):
|
with pytest.raises(TypeError, match=r"abstract"):
|
||||||
BIDSDataGrabber(datadir='/tmp', types='wrong',
|
BaseDataGrabber(datadir="/tmp", types=["func"]) # type: ignore
|
||||||
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_BIDSDataladDataGrabber():
|
def test_BaseDataGrabber() -> None:
|
||||||
"""Test BIDSDataladDataGrabber"""
|
"""Test BaseDataGrabber."""
|
||||||
types = ['T1w', 'bold']
|
# Create concrete class.
|
||||||
patterns = {
|
class MyDataGrabber(BaseDataGrabber):
|
||||||
'T1w': 'anat/{subject}_T1w.nii.gz',
|
def __getitem__(self, element):
|
||||||
'bold': 'func/{subject}_task-rest_bold.nii.gz'
|
return super().__getitem__(element)
|
||||||
}
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match=r"uri must be provided"):
|
def get_elements(self):
|
||||||
BIDSDataladDataGrabber(datadir=None, types=types, patterns=patterns)
|
return super().get_elements()
|
||||||
|
|
||||||
repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids'
|
dg = MyDataGrabber(datadir="/tmp", types=["func"])
|
||||||
rootdir = 'example_bids'
|
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,
|
with pytest.raises(NotImplementedError):
|
||||||
types=types, patterns=patterns) as dg:
|
dg.get_elements()
|
||||||
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:
|
with dg:
|
||||||
t_sub = dg[elem]
|
assert dg.datadir == Path("/tmp")
|
||||||
assert 'path' in t_sub['T1w']
|
assert dg.types == ["func"]
|
||||||
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'
|
|
||||||
|
|
|
||||||
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>
|
# 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>
|
||||||
# License: AGPL
|
# 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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from ..utils.logging import logger
|
from ..pipeline.pipeline_mixin import PipelineStepMixin
|
||||||
from ..markers.base 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 = {
|
_extensions = {
|
||||||
'.nii': 'NIFTI',
|
".nii": "NIFTI",
|
||||||
'.nii.gz': 'NIFTI',
|
".nii.gz": "NIFTI",
|
||||||
'.csv': 'CSV',
|
".csv": "CSV",
|
||||||
'.tsv': 'TSV'
|
".tsv": "TSV",
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
# Map each kind to a function and arguments
|
# Map each kind to a function and arguments
|
||||||
_readers = {}
|
_readers = {}
|
||||||
_readers['NIFTI'] = dict(func=nib.load, params=None)
|
_readers["NIFTI"] = {"func": nib.load, "params": None}
|
||||||
_readers['CSV'] = dict(func=pd.read_csv, params=None)
|
_readers["CSV"] = {"func": pd.read_csv, "params": None}
|
||||||
_readers['TSV'] = dict(func=pd.read_csv, params={'sep': '\t'})
|
_readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}}
|
||||||
|
|
||||||
|
|
||||||
class DefaultDataReader(PipelineStepMixin):
|
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
|
# Nothing to validate, any input is fine
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
# TODO: complete type annotations
|
||||||
def get_output_kind(self, input):
|
def get_output_kind(self, input):
|
||||||
|
"""Get output kind.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
input
|
||||||
|
|
||||||
|
"""
|
||||||
# It will output the same kind of data as the input
|
# It will output the same kind of data as the input
|
||||||
return 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
|
# 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:
|
if params is None:
|
||||||
params = {}
|
params = {}
|
||||||
for kind in input.keys():
|
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, {})
|
t_params = params.get(kind, {})
|
||||||
|
|
||||||
|
# Convert to Path if datareader is not well done
|
||||||
if not isinstance(t_path, Path):
|
if not isinstance(t_path, Path):
|
||||||
t_path = Path(t_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
|
fread = None
|
||||||
|
|
||||||
fname = t_path.name.lower()
|
fname = t_path.name.lower()
|
||||||
for ext, ftype in _extensions.items():
|
for ext, ftype in _extensions.items():
|
||||||
if fname.endswith(ext):
|
if fname.endswith(ext):
|
||||||
logger.info(f'Reading {ftype} file {t_path.as_posix()}')
|
logger.info(f"{kind} is type {ftype}")
|
||||||
reader_func = _readers[ftype]['func']
|
reader_func = _readers[ftype]["func"]
|
||||||
reader_params = _readers[ftype]['params']
|
reader_params = _readers[ftype]["params"]
|
||||||
if reader_params is not None:
|
if reader_params is not None:
|
||||||
t_params.update(reader_params)
|
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)
|
fread = reader_func(t_path, **t_params)
|
||||||
break
|
break
|
||||||
if fread is None:
|
if fread is None:
|
||||||
logger.info(
|
logger.info(
|
||||||
f'Unknown file type {t_path.as_posix()}, skipping reading')
|
f"Unknown file type {t_path.as_posix()}, skipping reading"
|
||||||
out[kind]['data'] = fread
|
)
|
||||||
|
out[kind]["data"] = fread
|
||||||
|
if "meta" not in out:
|
||||||
|
out["meta"] = {}
|
||||||
|
out["meta"]["datareader"] = self.get_meta()
|
||||||
return out
|
return out
|
||||||
|
|
|
||||||
|
|
@ -1,128 +1,168 @@
|
||||||
|
"""Provide tests for default data reader."""
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import tempfile
|
|
||||||
from numpy.testing import assert_array_equal
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
from nibabel import testing as nib_testing
|
|
||||||
import pandas as pd
|
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 pandas.testing import assert_frame_equal
|
||||||
|
|
||||||
from junifer.datareader import DefaultDataReader
|
from junifer.datareader import DefaultDataReader
|
||||||
|
|
||||||
|
|
||||||
def test_validation():
|
@pytest.mark.parametrize(
|
||||||
"""Test validating input/output"""
|
"kind", [["T1w", "BOLD", "T2", "dwi"], [], None, ["whatever"]]
|
||||||
kinds = [
|
)
|
||||||
['T1w', 'BOLD', 'T2', 'dwi'],
|
def test_validation(kind) -> None:
|
||||||
[],
|
"""Test validating input/output.
|
||||||
None,
|
|
||||||
['whatever']
|
|
||||||
]
|
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
kind : list of str or str or None
|
||||||
|
The parametrized kind of data.
|
||||||
|
|
||||||
|
"""
|
||||||
reader = DefaultDataReader()
|
reader = DefaultDataReader()
|
||||||
|
assert reader.validate_input(kind) is None
|
||||||
for t_kind in kinds:
|
assert reader.get_output_kind(kind) == kind
|
||||||
assert reader.validate_input(t_kind) is None
|
assert reader.validate(kind) == kind
|
||||||
assert reader.get_output_kind(t_kind) == t_kind
|
|
||||||
assert reader.validate(t_kind) == t_kind
|
|
||||||
|
|
||||||
|
|
||||||
def test_read_nifti():
|
def test_meta() -> None:
|
||||||
"""Test reading NIFTI files"""
|
"""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()
|
reader = DefaultDataReader()
|
||||||
nib_data_path = Path(nib_testing.data_path)
|
nib_data_path = Path(nib_testing.data_path)
|
||||||
|
|
||||||
for fname in ['example4d.nii.gz',
|
t_path = nib_data_path / fname
|
||||||
'reoriented_anat_moved.nii']:
|
|
||||||
t_path = nib_data_path / fname
|
|
||||||
|
|
||||||
input = {'bold': t_path}
|
input = {"bold": {"path": 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}
|
|
||||||
output = reader.fit_transform(input)
|
output = reader.fit_transform(input)
|
||||||
|
|
||||||
assert isinstance(output, dict)
|
assert isinstance(output, dict)
|
||||||
assert 'anat' in output
|
assert "bold" in output
|
||||||
assert isinstance(output['anat'], dict)
|
assert isinstance(output["bold"], dict)
|
||||||
assert 'path' in output['anat']
|
assert "path" in output["bold"]
|
||||||
assert isinstance(output['anat']['path'], Path)
|
assert "data" in output["bold"]
|
||||||
assert 'data' in output['anat']
|
|
||||||
assert output['anat']['data'] is not None
|
|
||||||
|
|
||||||
assert isinstance(output['whatever'], dict)
|
read_img = output["bold"]["data"]
|
||||||
assert 'path' in output['whatever']
|
|
||||||
assert isinstance(output['whatever']['path'], Path)
|
t_read_img = nib.load(t_path)
|
||||||
assert 'data' in output['whatever']
|
assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata())
|
||||||
assert output['whatever']['data'] is None
|
|
||||||
|
input = {"bold": {"path": t_path.as_posix()}}
|
||||||
|
output2 = reader.fit_transform(input)
|
||||||
|
assert output["bold"]["path"] == output2["bold"]["path"]
|
||||||
|
|
||||||
|
|
||||||
def test_read_csv():
|
def test_read_unknown() -> None:
|
||||||
"""Test reading CSV files"""
|
"""Test (not) reading unknown files."""
|
||||||
d = {'col1': [1, 2, 3, 4, 5], 'col2': [3, 4, 5, 6, 7]}
|
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)
|
df = pd.DataFrame(d)
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
|
||||||
tmpdir = Path(tmpdir)
|
|
||||||
df.to_csv(tmpdir / 'test.csv')
|
|
||||||
|
|
||||||
reader = DefaultDataReader()
|
df.to_csv(tmp_path / "test_read_csv.csv")
|
||||||
input = {'csv': tmpdir / 'test.csv'}
|
|
||||||
output = reader.fit_transform(input)
|
|
||||||
|
|
||||||
assert isinstance(output, dict)
|
reader = DefaultDataReader()
|
||||||
assert 'csv' in output
|
input = {"csv": {"path": tmp_path / "test_read_csv.csv"}}
|
||||||
assert isinstance(output['csv'], dict)
|
output = reader.fit_transform(input)
|
||||||
assert 'path' in output['csv']
|
|
||||||
assert 'data' in output['csv']
|
|
||||||
|
|
||||||
read_df = output['csv']['data'][['col1', 'col2']]
|
assert isinstance(output, dict)
|
||||||
assert_frame_equal(df, read_df)
|
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=';')
|
read_df = output["csv"]["data"][["col1", "col2"]]
|
||||||
input = {'csv': tmpdir / 'test.csv'}
|
assert_frame_equal(df, read_df)
|
||||||
params = {'csv': {'sep': ';'}}
|
|
||||||
output = reader.fit_transform(input, params)
|
|
||||||
|
|
||||||
assert isinstance(output, dict)
|
df.to_csv(tmp_path / "test_read_csv.csv", sep=";")
|
||||||
assert 'csv' in output
|
input = {"csv": {"path": tmp_path / "test_read_csv.csv"}}
|
||||||
assert isinstance(output['csv'], dict)
|
params = {"csv": {"sep": ";"}}
|
||||||
assert 'path' in output['csv']
|
output = reader.fit_transform(input, params)
|
||||||
assert 'data' in output['csv']
|
|
||||||
|
|
||||||
read_df = output['csv']['data'][['col1', 'col2']]
|
assert isinstance(output, dict)
|
||||||
assert_frame_equal(df, read_df)
|
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')
|
read_df = output["csv"]["data"][["col1", "col2"]]
|
||||||
input = {'csv': tmpdir / 'test.tsv'}
|
assert_frame_equal(df, read_df)
|
||||||
output = reader.fit_transform(input)
|
|
||||||
|
|
||||||
assert isinstance(output, dict)
|
df.to_csv(tmp_path / "test_read_csv.tsv", sep="\t")
|
||||||
assert 'csv' in output
|
input = {"csv": {"path": tmp_path / "test_read_csv.tsv"}}
|
||||||
assert isinstance(output['csv'], dict)
|
output = reader.fit_transform(input)
|
||||||
assert 'path' in output['csv']
|
|
||||||
assert 'data' in output['csv']
|
|
||||||
|
|
||||||
read_df = output['csv']['data'][['col1', 'col2']]
|
assert isinstance(output, dict)
|
||||||
assert_frame_equal(df, read_df)
|
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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Leonard Sasse <l.sasse@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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
class PipelineStepMixin():
|
from typing import Dict, List, Optional, Union
|
||||||
|
|
||||||
@property
|
from ..pipeline.pipeline_mixin import PipelineStepMixin
|
||||||
def name(self):
|
from ..utils import logger, raise_error
|
||||||
return self.__class__.__name__
|
|
||||||
|
|
||||||
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
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Junifer Data dictionary
|
kind : str
|
||||||
The input to the pipeline step.
|
The kind of 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.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
output : Junifer Data dictionary
|
dict
|
||||||
The output of the pipeline step.
|
The metadata as a dictionary.
|
||||||
"""
|
|
||||||
raise NotImplementedError('get_output_kind not implemented')
|
|
||||||
|
|
||||||
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
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Junifer Data dictionary
|
input : list of str
|
||||||
The input to the pipeline step.
|
The input to the pipeline step. The list must contain the
|
||||||
|
available Junifer Data dictionary keys.
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
output : Junifer Data dictionary
|
|
||||||
The output of the pipeline step.
|
|
||||||
|
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
ValueError:
|
ValueError
|
||||||
If the input does not have the required data.
|
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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# 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 ..utils import logger
|
||||||
from ..datareader import DefaultDataReader
|
|
||||||
|
|
||||||
|
|
||||||
class MarkerCollection():
|
class MarkerCollection:
|
||||||
def __init__(self, markers, datareader=None, preprocessing=None,
|
"""Class for marker collection.
|
||||||
storage=None):
|
|
||||||
|
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:
|
if datareader is None:
|
||||||
datareader = DefaultDataReader()
|
datareader = DefaultDataReader()
|
||||||
self._datareader = datareader
|
self._datareader = datareader
|
||||||
self._preprocessing = preprocessing
|
self._preprocessing = preprocessing
|
||||||
self._markers = markers
|
|
||||||
self._storage = storage
|
self._storage = storage
|
||||||
|
|
||||||
def fit(self, input):
|
def fit(self, input: Dict[str, Dict]) -> Optional[Dict]:
|
||||||
"""Fit the pipeline.
|
"""Fit the pipeline.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
input : Junifer Data dictionary (input)
|
input
|
||||||
The input data to fit the pipeline on. Should be the output of
|
The input data to fit the pipeline on. Should be the output of
|
||||||
indexing the DataGrabber with one element.
|
indexing the DataGrabber with one element.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
output : dict[str -> object]
|
output : dict or None
|
||||||
The output of the pipeline. Each key represents a marker name and
|
The output of the pipeline. Each key represents a marker name and
|
||||||
the values are the computer marker values. If the pipeline has a
|
the values are the computer marker values. If the pipeline has a
|
||||||
storage configured, then the output will be None.
|
storage configured, then the output will be None.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.info('Fitting pipeline')
|
logger.info("Fitting pipeline")
|
||||||
data = self._datareader.fit_transform(input)
|
data = self._datareader.fit_transform(input)
|
||||||
if self._preprocessing is not None:
|
if self._preprocessing is not None:
|
||||||
logger.info('Preprocessing data')
|
logger.info("Preprocessing data")
|
||||||
data = self._preprocessing.fit_transform(data)
|
data = self._preprocessing.fit_transform(data)
|
||||||
out = {}
|
out = {}
|
||||||
for marker in self._markers:
|
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)
|
m_value = marker.fit_transform(data, storage=self._storage)
|
||||||
if self._storage is None:
|
if self._storage is None:
|
||||||
out[marker.name] = m_value
|
out[marker.name] = m_value
|
||||||
|
logger.info("Marker collection fitting done")
|
||||||
return None if self._storage else out
|
return None if self._storage else out
|
||||||
|
|
||||||
def validate(self, datagrabber):
|
# TODO: complete type annotations
|
||||||
|
def validate(self, datagrabber) -> None:
|
||||||
"""Validate the pipeline.
|
"""Validate the pipeline.
|
||||||
|
|
||||||
Without doing any computation, check if the Marker Collection can
|
Without doing any computation, check if the Marker Collection can
|
||||||
be fit without problems. That is, the data required for each marker is
|
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,
|
present and streamed down the steps. Also, if a storage is configured,
|
||||||
check that the storage can handle the markers output.
|
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)
|
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:
|
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)
|
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:
|
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)
|
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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Leonard Sasse <l.sasse@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>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Leonard Sasse <l.sasse@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# 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