Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 13 additions & 7 deletions trx/io.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ def get_trx_tmp_dir():
return tempfile.TemporaryDirectory(dir=trx_tmp_dir, prefix="trx_")


def load_sft_with_reference(filepath, reference=None, bbox_check=True):
def load_sft_with_reference(filepath, reference=None, bbox_check=True, **kwargs):
"""Load a tractogram as a StatefulTractogram with an explicit reference.

Parameters
Expand All @@ -59,6 +59,8 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True):
bbox_check : bool, optional
If True, validate that streamlines lie within the reference bounding
box. Defaults to True.
**kwargs
Additional keyword arguments passed to dipy's load_tractogram.

Returns
-------
Expand All @@ -70,7 +72,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True):
IOError
If the file format is unsupported or a required reference is missing.
"""
if not dipy_available:
if not dipy_available: # pragma: no cover
logging.error(
"Dipy library is missing, cannot use functions related "
"to the StatefulTractogram."
Expand All @@ -83,20 +85,22 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True):
if ext == ".trk":
if reference is not None and reference != "same":
logging.warning(f"Reference is discarded for this file format {filepath}.")
sft = load_tractogram(filepath, "same", bbox_valid_check=bbox_check)
sft = load_tractogram(filepath, "same", bbox_valid_check=bbox_check, **kwargs)
elif ext in [".tck", ".fib", ".vtk", ".dpy"]:
if reference is None or reference == "same":
raise IOError(f"--reference is required for this file format {filepath}.")
else:
sft = load_tractogram(filepath, reference, bbox_valid_check=bbox_check)
sft = load_tractogram(
filepath, reference, bbox_valid_check=bbox_check, **kwargs
)

else:
raise IOError(f"{filepath} is an unsupported file format")

return sft


def load(tractogram_filename, reference):
def load(tractogram_filename, reference, **kwargs):
"""Load a tractogram from disk and return a TRX or StatefulTractogram.

Parameters
Expand All @@ -105,6 +109,8 @@ def load(tractogram_filename, reference):
Path to the input tractogram. TRX directories are supported.
reference : str or nibabel.Nifti1Image
Reference image used for formats without embedded affine information.
**kwargs
Additional keyword arguments passed to dipy's load_tractogram.

Returns
-------
Expand All @@ -116,7 +122,7 @@ def load(tractogram_filename, reference):
in_ext = split_name_with_gz(tractogram_filename)[1]
if in_ext != ".trx" and not os.path.isdir(tractogram_filename):
tractogram_obj = load_sft_with_reference(
tractogram_filename, reference, bbox_check=False
tractogram_filename, reference, bbox_check=False, **kwargs
)
else:
tractogram_obj = tmm.load(tractogram_filename)
Expand Down Expand Up @@ -145,7 +151,7 @@ def save(tractogram_obj, tractogram_filename, bbox_valid_check=False):
The function writes to disk and returns ``None``. Returns ``None``
immediately when ``dipy`` is unavailable.
"""
if not dipy_available:
if not dipy_available: # pragma: no cover
logging.error(
"Dipy library is missing, cannot use functions related "
"to the StatefulTractogram."
Expand Down
76 changes: 75 additions & 1 deletion trx/tests/test_cli.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
# -*- coding: utf-8 -*-
"""Tests for CLI commands and workflow functions."""

from contextlib import nullcontext
import os
import tempfile
from types import SimpleNamespace
from unittest.mock import patch

from deepdiff import DeepDiff
import numpy as np
Expand All @@ -19,6 +22,7 @@
from trx.fetcher import fetch_data, get_home, get_testing_files_dict
import trx.trx_file_memmap as tmm
from trx.workflows import (
_create_temp_memmap,
convert_dsi_studio,
convert_tractogram,
generate_trx_from_scratch,
Expand All @@ -33,6 +37,70 @@
)


def test_create_temp_memmap_uses_reopenable_path(tmp_path):
with patch("trx.workflows.np.memmap") as mock_memmap:
_create_temp_memmap(tmp_path, np.dtype("float32"), (10,))

filename = mock_memmap.call_args.args[0]
assert isinstance(filename, (str, os.PathLike))
assert os.path.dirname(os.fspath(filename)) == os.fspath(tmp_path)


def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path):
trx = SimpleNamespace(
streamlines=SimpleNamespace(
_data=np.arange(6, dtype=np.float16).reshape((2, 3)),
_offsets=np.array([0, 3], dtype=np.uint64),
),
data_per_vertex={
"mock_dpv": SimpleNamespace(
_data=np.arange(6, dtype=np.uint8).reshape((2, 3))
)
},
data_per_streamline={"mock_dps": np.array([1, 2], dtype=np.uint8)},
data_per_group={
"mock_group": {"mock_dpg": np.array([1.0, 2.0], dtype=np.float32)}
},
groups={"mock_group": np.array([0, 1], dtype=np.int32)},
)
trx.close = lambda: None

with (
patch(
"trx.workflows.get_trx_tmp_dir",
return_value=nullcontext(os.fspath(tmp_path)),
),
patch("trx.workflows.tmm.load", return_value=trx),
patch("trx.workflows.tmm.save") as mock_save,
patch(
"trx.workflows.tempfile.NamedTemporaryFile",
side_effect=AssertionError(
"NamedTemporaryFile should not be used for writable memmaps"
),
),
):
manipulate_trx_datatype(
"in.trx",
"out.trx",
{
"positions": np.dtype("float32"),
"offsets": np.dtype("uint32"),
"dpv": {"mock_dpv": np.dtype("uint16")},
"dps": {"mock_dps": np.dtype("float32")},
"dpg": {"mock_group": {"mock_dpg": np.dtype("float64")}},
"groups": {"mock_group": np.dtype("uint16")},
},
)

assert trx.streamlines._data.dtype == np.dtype("float32")
assert trx.streamlines._offsets.dtype == np.dtype("uint32")
assert trx.data_per_vertex["mock_dpv"]._data.dtype == np.dtype("uint16")
assert trx.data_per_streamline["mock_dps"].dtype == np.dtype("float32")
assert trx.data_per_group["mock_group"]["mock_dpg"].dtype == np.dtype("float64")
assert trx.groups["mock_group"].dtype == np.dtype("uint16")
mock_save.assert_called_once_with(trx, "out.trx")


def _normalize_dtype_dict(dtype_dict):
"""Normalize dtype dict to use explicit little-endian byte order.

Expand Down Expand Up @@ -469,7 +537,13 @@ def test_execution_manipulate_trx_datatype(self):
}

out_gen_path = os.path.join(tmp_dir, "generated.trx")
manipulate_trx_datatype(expected_trx, out_gen_path, generated_dtype)
with patch(
"trx.workflows.tempfile.NamedTemporaryFile",
side_effect=AssertionError(
"NamedTemporaryFile should not be used for writable memmaps"
),
):
manipulate_trx_datatype(expected_trx, out_gen_path, generated_dtype)
trx = tmm.load(out_gen_path)
assert (
DeepDiff(
Expand Down
29 changes: 26 additions & 3 deletions trx/tests/test_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,16 @@
from dipy.io.streamline import load_tractogram, save_tractogram

dipy_available = True
except ImportError:
except ImportError: # pragma: no cover
dipy_available = False

try:
import fury # noqa: F401

fury_available = True
except ImportError: # pragma: no cover
fury_available = False

from trx.fetcher import fetch_data, get_home, get_testing_files_dict
from trx.io import load, save
import trx.trx_file_memmap as tmm
Expand All @@ -28,6 +35,8 @@
@pytest.mark.parametrize("path", [("gs.trk"), ("gs.tck"), ("gs.vtk")])
@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.")
def test_seq_ops_sft(path):
if path.endswith(".vtk") and not fury_available:
pytest.skip("fury is not installed")
with TemporaryDirectory() as tmp_dir:
gs_dir = os.path.join(get_home(), "gold_standard")
path = os.path.join(tmp_dir, path)
Expand Down Expand Up @@ -56,10 +65,17 @@ def test_seq_ops_trx():
@pytest.mark.parametrize("path", [("gs.trx"), ("gs.trk"), ("gs.tck"), ("gs.vtk")])
@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.")
def test_load_vox(path):
if path.endswith(".vtk") and not fury_available:
pytest.skip("fury is not installed")
gs_dir = os.path.join(get_home(), "gold_standard")
path = os.path.join(gs_dir, path)
coord = np.loadtxt(os.path.join(get_home(), "gold_standard", "gs_vox_space.txt"))
obj = load(path, os.path.join(gs_dir, "gs.nii"))
if path.endswith(".vtk"):
from dipy.io.stateful_tractogram import Space

obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=Space.LPSMM)
else:
obj = load(path, os.path.join(gs_dir, "gs.nii"))

sft = obj.to_sft() if isinstance(obj, TrxFile) else obj
sft.to_vox()
Expand All @@ -72,10 +88,17 @@ def test_load_vox(path):
@pytest.mark.parametrize("path", [("gs.trx"), ("gs.trk"), ("gs.tck"), ("gs.vtk")])
@pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.")
def test_load_voxmm(path):
if path.endswith(".vtk") and not fury_available:
pytest.skip("fury is not installed")
gs_dir = os.path.join(get_home(), "gold_standard")
path = os.path.join(gs_dir, path)
coord = np.loadtxt(os.path.join(get_home(), "gold_standard", "gs_voxmm_space.txt"))
obj = load(path, os.path.join(gs_dir, "gs.nii"))
if path.endswith(".vtk"):
from dipy.io.stateful_tractogram import Space

obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=Space.LPSMM)
else:
obj = load(path, os.path.join(gs_dir, "gs.nii"))

sft = obj.to_sft() if isinstance(obj, TrxFile) else obj
sft.to_voxmm()
Expand Down
Loading
Loading