From e511ab318d29daee1c84a7cd12b32f39c7a9eeaa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jon=20Haitz=20Legarreta=20Gorro=C3=B1o?= Date: Sat, 22 Aug 2026 14:34:09 +0200 Subject: [PATCH] REF: Adopt `pathlib` for path manipulation Adopt `pathlib` for path manipulation: - Use `pathlib` for path manipulation in all internal operations, including using `PurePosixPath` instance for ZIP member names, as they have POSIX-style path semantics. - Simplifies the code and eases the manipulation of paths. `pathlib` documentation: https://docs.python.org/3/library/pathlib.html --- .spin/cmds.py | 46 +++---- docs/source/conf.py | 3 +- examples/plot_dps_dpv.py | 4 +- examples/plot_groups.py | 4 +- examples/plot_read_write_trx.py | 7 +- trx/fetcher.py | 15 +- trx/io.py | 20 ++- trx/tests/test_cli.py | 96 ++++++------- trx/tests/test_io.py | 125 +++++++++-------- trx/tests/test_memmap.py | 48 +++---- trx/tests/test_utils.py | 3 +- trx/trx_file_memmap.py | 236 ++++++++++++++++++-------------- trx/utils.py | 23 ++-- trx/workflows.py | 28 ++-- 14 files changed, 346 insertions(+), 312 deletions(-) diff --git a/.spin/cmds.py b/.spin/cmds.py index d4fcd27e..e6810104 100644 --- a/.spin/cmds.py +++ b/.spin/cmds.py @@ -1,7 +1,7 @@ """Custom spin commands for trx-python development.""" -import glob import os +from pathlib import Path import shutil import subprocess import sys @@ -223,35 +223,32 @@ def docs(_clean, open_browser): open_browser : bool If True, open documentation in browser after building. """ - import os - docs_dir = "docs" + docs_dir = Path("docs") if _clean: click.echo("Cleaning build directory...") - build_dir = os.path.join(docs_dir, "_build") - if os.path.exists(build_dir): + build_dir = docs_dir / "_build" + if build_dir.exists(): shutil.rmtree(build_dir) # Clean sphinx-gallery generated files - gallery_dir = os.path.join(docs_dir, "source", "auto_examples") - if os.path.exists(gallery_dir): + gallery_dir = docs_dir / "source" / "auto_examples" + if gallery_dir.exists(): click.echo("Cleaning sphinx-gallery generated files...") shutil.rmtree(gallery_dir) # Clean sphinx-gallery execution times file - sg_times = os.path.join(docs_dir, "source", "sg_execution_times.rst") - if os.path.exists(sg_times): + sg_times = docs_dir / "source" / "sg_execution_times.rst" + if sg_times.exists(): os.remove(sg_times) click.echo("Building documentation...") - cmd = ["make", "-C", docs_dir, "html"] + cmd = ["make", "-C", str(docs_dir), "html"] result = run(cmd, capture=False, check=False) if result == 0: - index_path = os.path.abspath( - os.path.join(docs_dir, "_build", "html", "index.html") - ) + index_path = (docs_dir / "_build" / "html" / "index.html").resolve() click.echo("\nDocs built successfully!") click.echo(f"Open: {index_path}") @@ -270,27 +267,26 @@ def clean(): # noqa: C901 # Clean TRX temp directory trx_tmp_dir = os.getenv("TRX_TMPDIR", tempfile.gettempdir()) - if os.path.exists(trx_tmp_dir): - temp_files = glob.glob(os.path.join(trx_tmp_dir, "trx_*")) - for temp_dir in temp_files: - if os.path.isdir(temp_dir): - click.echo(f"Removing temporary directory: {temp_dir}") - shutil.rmtree(temp_dir) + if trx_tmp_dir.exists(): + for temp_name in trx_tmp_dir.glob("trx_*"): + if temp_name.is_dir(): + click.echo(f"Removing temporary directory: {temp_name}") + shutil.rmtree(temp_name) # Clean build artifacts for build_pattern in ["build", "dist", "*.egg-info"]: - for path in glob.glob(build_pattern): - if os.path.isdir(path): + for path in Path(".").glob(build_pattern): + if path.is_dir(): click.echo(f"Removing build directory: {path}") shutil.rmtree(path) - elif os.path.isfile(path): + elif path.is_file(): click.echo(f"Removing build file: {path}") - os.remove(path) + path.unlink() # Clean Python cache for cache_dir in ["**/__pycache__", "**/.pytest_cache"]: - for path in glob.glob(cache_dir, recursive=True): - if os.path.isdir(path): + for path in Path(".").glob(cache_dir): + if path.is_dir(): click.echo(f"Removing cache directory: {path}") shutil.rmtree(path) diff --git a/docs/source/conf.py b/docs/source/conf.py index 417dc367..0faf990a 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -17,6 +17,7 @@ import warnings import os from datetime import datetime as dt +from pathlib import Path # -- Version information ----------------------------------------------------- # Get version from environment variable (set by CI) or package @@ -191,5 +192,5 @@ def _validate_reference_urls(urls, timeout=5): 'gallery_dirs': 'auto_examples', 'within_subsection_order': 'NumberOfCodeLinesSortKey', 'reference_url': _validate_reference_urls(_reference_urls), - 'default_thumb_file': os.path.join(os.path.dirname(__file__), '..', '_static', 'trx_logo.png'), + 'default_thumb_file': str(Path(__file__).parent / '..' / '_static' / 'trx_logo.png'), } diff --git a/examples/plot_dps_dpv.py b/examples/plot_dps_dpv.py index f6a6bedf..c352e558 100644 --- a/examples/plot_dps_dpv.py +++ b/examples/plot_dps_dpv.py @@ -39,8 +39,6 @@ # # Let's load a TRX file and explore its metadata. -import os - import numpy as np from trx.fetcher import fetch_data, get_home, get_testing_files_dict @@ -49,7 +47,7 @@ # Download test data fetch_data(get_testing_files_dict(), keys="gold_standard.zip") trx_home = get_home() -trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") +trx_path = trx_home / "gold_standard" / "gs.trx" # Load the TRX file tgm = load(trx_path) diff --git a/examples/plot_groups.py b/examples/plot_groups.py index 92eb9801..2196e13a 100644 --- a/examples/plot_groups.py +++ b/examples/plot_groups.py @@ -35,8 +35,6 @@ # # Let's load a TRX file that contains group information. -import os - import numpy as np from trx.fetcher import fetch_data, get_home, get_testing_files_dict @@ -45,7 +43,7 @@ # Download test data fetch_data(get_testing_files_dict(), keys="gold_standard.zip") trx_home = get_home() -trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") +trx_path = trx_home / "gold_standard" / "gs.trx" # Load the TRX file tgm = load(trx_path) diff --git a/examples/plot_read_write_trx.py b/examples/plot_read_write_trx.py index a9fa10d7..9a5f1790 100644 --- a/examples/plot_read_write_trx.py +++ b/examples/plot_read_write_trx.py @@ -24,7 +24,6 @@ # Let's start by loading an existing TRX file. First, we need to download # some test data. -import os import tempfile from trx.fetcher import fetch_data, get_home, get_testing_files_dict @@ -33,7 +32,7 @@ # Download test data fetch_data(get_testing_files_dict(), keys="gold_standard.zip") trx_home = get_home() -trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") +trx_path = trx_home / "gold_standard" / "gs.trx" # Load the TRX file tgm = load(trx_path) @@ -109,10 +108,10 @@ with tempfile.TemporaryDirectory() as tmpdir: # Save as TRX file (zip archive) - output_path = os.path.join(tmpdir, "output.trx") + output_path = tmpdir / "output.trx" save(tgm, output_path) print(f"Saved TRX file to: {output_path}") - print(f"File size: {os.path.getsize(output_path)} bytes") + print(f"File size: {output_path.stat().st_size} bytes") # Reload to verify reloaded = load(output_path) diff --git a/trx/fetcher.py b/trx/fetcher.py index 1323c852..e0b5a4da 100644 --- a/trx/fetcher.py +++ b/trx/fetcher.py @@ -4,6 +4,7 @@ import hashlib import logging import os +from pathlib import Path import shutil import urllib.request @@ -24,13 +25,13 @@ def get_home(): Returns ------- - str + Path Path to the TRX home directory. """ if "TRX_HOME" in os.environ: - trx_home = os.environ["TRX_HOME"] + trx_home = Path(os.environ["TRX_HOME"]) else: - trx_home = os.path.join(os.path.expanduser("~"), ".tee_ar_ex") + trx_home = Path("~").expanduser() / ".tee_ar_ex" return trx_home @@ -132,7 +133,7 @@ def fetch_data(files_dict, keys=None): # noqa: C901 """ trx_home = get_home() - if not os.path.exists(trx_home): + if not trx_home.exists(): os.makedirs(trx_home) if keys is None: @@ -147,10 +148,10 @@ def fetch_data(files_dict, keys=None): # noqa: C901 expected_sha = None else: url, expected_md5, expected_sha = file_entry - full_path = os.path.join(trx_home, fname) + full_path = trx_home / fname logging.info(f"Downloading {fname} to {trx_home}") - if not os.path.exists(full_path): + if not full_path.exists(): urllib.request.urlretrieve(url, full_path) actual_md5 = md5sum(full_path) @@ -169,5 +170,5 @@ def fetch_data(files_dict, keys=None): # noqa: C901 ) if fname.endswith(".zip"): - dst_dir = os.path.join(trx_home, fname[:-4]) + dst_dir = trx_home / fname[:-4] shutil.unpack_archive(full_path, extract_dir=dst_dir, format="zip") diff --git a/trx/io.py b/trx/io.py index 1fdc787b..99b52a41 100644 --- a/trx/io.py +++ b/trx/io.py @@ -3,6 +3,7 @@ import logging import os +from pathlib import Path import sys import tempfile @@ -32,18 +33,22 @@ def get_trx_tmp_dir(): """ if os.getenv("TRX_TMPDIR") is not None: if os.getenv("TRX_TMPDIR") == "use_working_dir": - trx_tmp_dir = os.getcwd() + trx_tmp_dir = str(Path.cwd()) else: trx_tmp_dir = os.getenv("TRX_TMPDIR") else: trx_tmp_dir = tempfile.gettempdir() if sys.version_info[1] >= 10: - return tempfile.TemporaryDirectory( + tmp_dir = tempfile.TemporaryDirectory( dir=trx_tmp_dir, prefix="trx_", ignore_cleanup_errors=True ) else: - return tempfile.TemporaryDirectory(dir=trx_tmp_dir, prefix="trx_") + tmp_dir = tempfile.TemporaryDirectory(dir=trx_tmp_dir, prefix="trx_") + + # Keep the TemporaryDirectory instance, but provide a Path-compatible name + tmp_dir.name = Path(tmp_dir.name) + return tmp_dir def load_sft_with_reference(filepath, reference=None, bbox_check=True, from_space=None): @@ -51,7 +56,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True, from_spac Parameters ---------- - filepath : str + filepath : str or Path Path to the tractogram file (.trk, .tck, .fib, .vtk, .dpy). reference : str or nibabel.Nifti1Image, optional Reference image used for formats without embedded affine information. @@ -81,7 +86,7 @@ def load_sft_with_reference(filepath, reference=None, bbox_check=True, from_spac from dipy.io.streamline import load_tractogram # Force the usage of --reference for all file formats without an header - _, ext = os.path.splitext(filepath) + ext = "".join(Path(filepath).suffixes) if ext == ".trk": if reference is not None and reference != "same": logging.warning(f"Reference is discarded for this file format {filepath}.") @@ -122,7 +127,7 @@ def load(tractogram_filename, reference=None, from_space=None): import trx.trx_file_memmap as tmm in_ext = split_name_with_gz(tractogram_filename)[1] - if in_ext != ".trx" and not os.path.isdir(tractogram_filename): + if in_ext != ".trx" and not tractogram_filename.is_dir(): tractogram_obj = load_sft_with_reference( tractogram_filename, reference, bbox_check=False, from_space=from_space ) @@ -140,7 +145,7 @@ def save(tractogram_obj, tractogram_filename, bbox_valid_check=False): tractogram_obj : TrxFile or StatefulTractogram Tractogram to persist. Non-TRX inputs are converted to StatefulTractogram before saving to non-TRX formats. - tractogram_filename : str + tractogram_filename : str or Path Destination file name. ``.trx`` will be saved using the TRX writer; all other extensions are handled by ``dipy.save_tractogram``. bbox_valid_check : bool, optional @@ -153,6 +158,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. """ + tractogram_filename = Path(tractogram_filename) if not dipy_available: logging.error( "Dipy library is missing, cannot use functions related " diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index d8fecfda..bc736502 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -3,6 +3,7 @@ from contextlib import nullcontext import os +from pathlib import Path from types import SimpleNamespace from unittest.mock import patch @@ -42,7 +43,7 @@ def test_create_temp_memmap_uses_reopenable_path(tmp_path): filename = mock_memmap.call_args.args[0] assert isinstance(filename, (str, os.PathLike)) - assert os.path.dirname(os.fspath(filename)) == os.fspath(tmp_path) + assert Path(filename).parent == tmp_path def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): @@ -67,7 +68,7 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): with ( patch( "trx.workflows.get_trx_tmp_dir", - return_value=nullcontext(os.fspath(tmp_path)), + return_value=nullcontext(str(tmp_path)), ), patch("trx.workflows.tmm.load", return_value=tgm), patch("trx.workflows.tmm.save") as mock_save, @@ -208,7 +209,7 @@ def test_trx_info_help(self, script_runner): def test_trx_info_execution(self, script_runner): """Test trx info command execution on a real TRX file.""" - trx_path = os.path.join(get_home(), "gold_standard", "gs.trx") + trx_path = get_home() / "gold_standard" / "gs.trx" ret = script_runner.run(["trx", "info", trx_path]) assert ret.success # Check key output elements @@ -220,7 +221,7 @@ def test_trx_info_execution(self, script_runner): def test_trx_info_wrong_extension(self, script_runner): """Test trx info rejects non-TRX files.""" - tck_path = os.path.join(get_home(), "gold_standard", "gs.tck") + tck_path = get_home() / "gold_standard" / "gs.tck" ret = script_runner.run(["trx", "info", tck_path]) assert not ret.success assert "not a TRX file" in ret.stderr @@ -238,11 +239,11 @@ class TestWorkflowFunctions: @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_dsi(self, tmp_path): - in_trk = os.path.join(get_home(), "DSI", "CC.trk.gz") - in_nii = os.path.join(get_home(), "DSI", "CC.nii.gz") - exp_data = os.path.join(get_home(), "DSI", "CC_fix_data.npy") - exp_offsets = os.path.join(get_home(), "DSI", "CC_fix_offsets.npy") - out_fix_path = os.path.join(tmp_path, "fixed.trk") + in_trk = get_home() / "DSI" / "CC.trk.gz" + in_nii = get_home() / "DSI" / "CC.nii.gz" + exp_data = get_home() / "DSI" / "CC_fix_data.npy" + exp_offsets = get_home() / "DSI" / "CC_fix_offsets.npy" + out_fix_path = tmp_path / "fixed.trk" convert_dsi_studio( in_trk, in_nii, out_fix_path, remove_invalid=False, keep_invalid=True ) @@ -256,10 +257,10 @@ def test_execution_convert_dsi(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_to_trx(self, tmp_path): - in_trk = os.path.join(get_home(), "DSI", "CC_fix.trk") - exp_data = os.path.join(get_home(), "DSI", "CC_fix_data.npy") - exp_offsets = os.path.join(get_home(), "DSI", "CC_fix_offsets.npy") - out_trx_path = os.path.join(tmp_path, "CC_fix.trx") + in_trk = get_home() / "DSI" / "CC_fix.trk" + exp_data = get_home() / "DSI" / "CC_fix_data.npy" + exp_offsets = get_home() / "DSI" / "CC_fix_offsets.npy" + out_trx_path = tmp_path / "CC_fix.trx" convert_tractogram(in_trk, out_trx_path, None) data_fix = np.load(exp_data) @@ -274,15 +275,15 @@ def test_execution_convert_to_trx(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_from_trx(self, tmp_path): - in_trk = os.path.join(get_home(), "DSI", "CC_fix.trk") - in_nii = os.path.join(get_home(), "DSI", "CC.nii.gz") - exp_data = os.path.join(get_home(), "DSI", "CC_fix_data.npy") - exp_offsets = os.path.join(get_home(), "DSI", "CC_fix_offsets.npy") + in_trk = get_home() / "DSI" / "CC_fix.trk" + in_nii = get_home() / "DSI" / "CC.nii.gz" + exp_data = get_home() / "DSI" / "CC_fix_data.npy" + exp_offsets = get_home() / "DSI" / "CC_fix_offsets.npy" # Sequential conversions - out_trx_path = os.path.join(tmp_path, "CC_fix.trx") - out_trk_path = os.path.join(tmp_path, "CC_fix.trk") - out_tck_path = os.path.join(tmp_path, "CC_fix.tck") + out_trx_path = tmp_path / "CC_fix.trx" + out_trk_path = tmp_path / "CC_fix.trk" + out_tck_path = tmp_path / "CC_fix.tck" convert_tractogram(in_trk, out_trx_path, None) convert_tractogram(out_trx_path, out_tck_path, None) convert_tractogram(out_trx_path, out_trk_path, None) @@ -300,8 +301,8 @@ def test_execution_convert_from_trx(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_dtype_p16_o64(self, tmp_path): - in_trk = os.path.join(get_home(), "DSI", "CC_fix.trk") - out_convert_path = os.path.join(tmp_path, "CC_fix_p16_o64.trx") + in_trk = get_home() / "DSI" / "CC_fix.trk" + out_convert_path = tmp_path / "CC_fix_p16_o64.trx" convert_tractogram( in_trk, out_convert_path, @@ -317,8 +318,8 @@ def test_execution_convert_dtype_p16_o64(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_dtype_p64_o32(self, tmp_path): - in_trk = os.path.join(get_home(), "DSI", "CC_fix.trk") - out_convert_path = os.path.join(tmp_path, "CC_fix_p16_o64.trx") + in_trk = get_home() / "DSI" / "CC_fix.trk" + out_convert_path = tmp_path / "CC_fix_p16_o64.trx" convert_tractogram( in_trk, out_convert_path, @@ -333,44 +334,44 @@ def test_execution_convert_dtype_p64_o32(self, tmp_path): tgm.close() def test_execution_generate_trx_from_scratch(self, tmp_path): - reference_fa = os.path.join(get_home(), "trx_from_scratch", "fa.nii.gz") - raw_arr_dir = os.path.join(get_home(), "trx_from_scratch", "test_npy") - expected_trx = os.path.join(get_home(), "trx_from_scratch", "expected.trx") + reference_fa = get_home() / "trx_from_scratch" / "fa.nii.gz" + raw_arr_dir = get_home() / "trx_from_scratch" / "test_npy" + expected_trx = get_home() / "trx_from_scratch" / "expected.trx" dpv = [ - (os.path.join(raw_arr_dir, "dpv_cx.npy"), "uint8"), - (os.path.join(raw_arr_dir, "dpv_cy.npy"), "uint8"), - (os.path.join(raw_arr_dir, "dpv_cz.npy"), "uint8"), + (raw_arr_dir / "dpv_cx.npy", "uint8"), + (raw_arr_dir / "dpv_cy.npy", "uint8"), + (raw_arr_dir / "dpv_cz.npy", "uint8"), ] dps = [ - (os.path.join(raw_arr_dir, "dps_algo.npy"), "uint8"), - (os.path.join(raw_arr_dir, "dps_cw.npy"), "float64"), + (raw_arr_dir / "dps_algo.npy", "uint8"), + (raw_arr_dir / "dps_cw.npy", "float64"), ] dpg = [ ( "g_AF_L", - os.path.join(raw_arr_dir, "dpg_AF_L_mean_fa.npy"), + raw_arr_dir / "dpg_AF_L_mean_fa.npy", "float32", ), ( "g_AF_R", - os.path.join(raw_arr_dir, "dpg_AF_R_mean_fa.npy"), + raw_arr_dir / "dpg_AF_R_mean_fa.npy", "float32", ), - ("g_AF_L", os.path.join(raw_arr_dir, "dpg_AF_L_volume.npy"), "float32"), + ("g_AF_L", raw_arr_dir / "dpg_AF_L_volume.npy", "float32"), ] groups = [ - (os.path.join(raw_arr_dir, "g_AF_L.npy"), "int32"), - (os.path.join(raw_arr_dir, "g_AF_R.npy"), "int32"), - (os.path.join(raw_arr_dir, "g_CST_L.npy"), "int32"), + (raw_arr_dir / "g_AF_L.npy", "int32"), + (raw_arr_dir / "g_AF_R.npy", "int32"), + (raw_arr_dir / "g_CST_L.npy", "int32"), ] - out_gen_path = os.path.join(tmp_path, "generated.trx") + out_gen_path = tmp_path / "generated.trx" generate_trx_from_scratch( reference_fa, out_gen_path, - positions=os.path.join(raw_arr_dir, "positions.npy"), - offsets=os.path.join(raw_arr_dir, "offsets.npy"), + positions=raw_arr_dir / "positions.npy", + offsets=raw_arr_dir / "offsets.npy", positions_dtype="float16", offsets_dtype="uint64", space_str="rasmm", @@ -419,8 +420,8 @@ def test_execution_generate_trx_from_scratch(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_concatenate_validate_trx(self, tmp_path): - trx1 = tmm.load(os.path.join(get_home(), "gold_standard", "gs.trx")) - trx2 = tmm.load(os.path.join(get_home(), "gold_standard", "gs.trx")) + trx1 = tmm.load(get_home() / "gold_standard" / "gs.trx") + trx2 = tmm.load(get_home() / "gold_standard" / "gs.trx") tgm = tmm.concatenate([trx1, trx2], preallocation=False) # Right size @@ -454,8 +455,8 @@ def test_execution_concatenate_validate_trx(self, tmp_path): ) # Validate - out_concat_path = os.path.join(tmp_path, "concat.trx") - out_valid_path = os.path.join(tmp_path, "valid.trx") + out_concat_path = tmp_path / "concat.trx" + out_valid_path = tmp_path / "valid.trx" tmm.save(tgm, out_concat_path) validate_tractogram( out_concat_path, @@ -477,7 +478,7 @@ def test_execution_concatenate_validate_trx(self, tmp_path): @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_manipulate_trx_datatype(self, tmp_path): - expected_trx = os.path.join(get_home(), "trx_from_scratch", "expected.trx") + expected_trx = get_home() / "trx_from_scratch" / "expected.trx" tgm = tmm.load(expected_trx) expected_dtype = { @@ -526,8 +527,7 @@ def test_execution_manipulate_trx_datatype(self, tmp_path): }, "groups": {"g_AF_L": np.dtype("uint16"), "g_AF_R": np.dtype("uint16")}, } - - out_gen_path = os.path.join(tmp_path, "generated.trx") + out_gen_path = tmp_path / "generated.trx" with patch( "trx.workflows.tempfile.NamedTemporaryFile", side_effect=AssertionError( diff --git a/trx/tests/test_io.py b/trx/tests/test_io.py index 4b0678a4..e07dba72 100644 --- a/trx/tests/test_io.py +++ b/trx/tests/test_io.py @@ -2,6 +2,7 @@ from copy import deepcopy import os +from pathlib import Path import tempfile import zipfile @@ -28,26 +29,26 @@ @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(tmp_path, path): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(tmp_path, path) + gs_dir = get_home() / "gold_standard" + path = tmp_path / path - obj = load(os.path.join(gs_dir, "gs.trx"), os.path.join(gs_dir, "gs.nii")) + obj = load(gs_dir / "gs.trx", gs_dir / "gs.nii") sft_1 = obj.to_sft() save_tractogram(sft_1, path) obj.close() - save_tractogram(sft_1, os.path.join(tmp_path, "tmp.trk")) + save_tractogram(sft_1, tmp_path / "tmp.trk") - _ = load_tractogram(os.path.join(tmp_path, "tmp.trk"), "same") + _ = load_tractogram(tmp_path / "tmp.trk", "same") def test_seq_ops_trx(tmp_path): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, "gs.trx") + gs_dir = get_home() / "gold_standard" + path = gs_dir / "gs.trx" trx_1 = tmm.load(path) - tmm.save(trx_1, os.path.join(tmp_path, "tmp.trx")) + tmm.save(trx_1, tmp_path / "tmp.trx") trx_1.close() - trx_2 = tmm.load(os.path.join(tmp_path, "tmp.trx")) + trx_2 = tmm.load(tmp_path / "tmp.trx") trx_2.close() @@ -56,11 +57,11 @@ def test_seq_ops_trx(tmp_path): def test_load_vox(path): from dipy.io.stateful_tractogram import Space - 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")) - from_space = Space.LPSMM if path.endswith("gs.vtk") else None - obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=from_space) + gs_dir = get_home() / "gold_standard" + path = gs_dir / path + coord = np.loadtxt(get_home() / "gold_standard" / "gs_vox_space.txt") + from_space = Space.LPSMM if path.name.endswith("gs.vtk") else None + obj = load(path, gs_dir / "gs.nii", from_space=from_space) sft = obj.to_sft() if isinstance(obj, TrxFile) else obj sft.to_vox() @@ -75,11 +76,11 @@ def test_load_vox(path): def test_load_voxmm(path): from dipy.io.stateful_tractogram import Space - 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")) - from_space = Space.LPSMM if path.endswith("gs.vtk") else None - obj = load(path, os.path.join(gs_dir, "gs.nii"), from_space=from_space) + gs_dir = get_home() / "gold_standard" + path = gs_dir / path + coord = np.loadtxt(get_home() / "gold_standard" / "gs_voxmm_space.txt") + from_space = Space.LPSMM if path.name.endswith("gs.vtk") else None + obj = load(path, gs_dir / "gs.nii", from_space=from_space) sft = obj.to_sft() if isinstance(obj, TrxFile) else obj sft.to_voxmm() @@ -92,20 +93,21 @@ def test_load_voxmm(path): @pytest.mark.parametrize("path", ["gs.trk", "gs.trx", "gs_fldr.trx"]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_multi_load_save_rasmm(tmp_path, path): - gs_dir = os.path.join(get_home(), "gold_standard") - basename, ext = os.path.splitext(path) + gs_dir = get_home() / "gold_standard" + basename = Path(path).stem + ext = Path(path).suffix.lstrip(".") - path = os.path.join(gs_dir, path) - coord = np.loadtxt(os.path.join(get_home(), "gold_standard", "gs_rasmm_space.txt")) + path = gs_dir / path + coord = np.loadtxt(get_home() / "gold_standard" / "gs_rasmm_space.txt") - obj = load(path, os.path.join(gs_dir, "gs.nii")) + obj = load(path, gs_dir / "gs.nii") for i in range(3): - out_path = os.path.join(tmp_path, f"{basename}_tmp{i}_{ext}") + out_path = tmp_path / f"{basename}_tmp{i}_{ext}" save(obj, out_path) if isinstance(obj, TrxFile): obj.close() - obj = load(out_path, os.path.join(gs_dir, "gs.nii")) + obj = load(out_path, gs_dir / "gs.nii") assert_allclose(obj.streamlines._data, coord, rtol=1e-04, atol=1e-06) if isinstance(obj, TrxFile): @@ -115,26 +117,22 @@ def test_multi_load_save_rasmm(tmp_path, path): @pytest.mark.parametrize("path", ["gs.trx", "gs_fldr.trx"]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_delete_tmp_gs_dir(path): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, path) + gs_dir = get_home() / "gold_standard" + path = gs_dir / path trx1 = tmm.load(path) - if os.path.isfile(path): + if path.is_file(): tmp_gs_dir = deepcopy(trx1._uncompressed_folder_handle.name) - assert os.path.isdir(tmp_gs_dir) + assert tmp_gs_dir.is_dir() sft = trx1.to_sft() trx1.close() - coord_rasmm = np.loadtxt( - os.path.join(get_home(), "gold_standard", "gs_rasmm_space.txt") - ) - coord_vox = np.loadtxt( - os.path.join(get_home(), "gold_standard", "gs_vox_space.txt") - ) + coord_rasmm = np.loadtxt(get_home() / "gold_standard" / "gs_rasmm_space.txt") + coord_vox = np.loadtxt(get_home() / "gold_standard" / "gs_vox_space.txt") # The folder trx representation does not need tmp files - if os.path.isfile(path): - assert not os.path.isdir(tmp_gs_dir) + if path.is_file(): + assert not tmp_gs_dir.is_dir() assert_allclose(sft.streamlines._data, coord_rasmm, rtol=1e-04, atol=1e-06) @@ -156,8 +154,8 @@ def test_delete_tmp_gs_dir(path): @pytest.mark.parametrize("path", ["gs.trx"]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_close_tmp_files(path): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, path) + gs_dir = get_home() / "gold_standard" + path = gs_dir / path tgm = tmm.load(path) process = psutil.Process(os.getpid()) @@ -175,7 +173,7 @@ def test_close_tmp_files(path): count = 0 for open_file in open_files: - basename = os.path.basename(open_file.path) + basename = Path(open_file.path).name if basename in expected_content: count += 1 @@ -185,7 +183,7 @@ def test_close_tmp_files(path): open_files = process.open_files() count = 0 for open_file in open_files: - basename = os.path.basename(open_file.path) + basename = Path(open_file.path).name if basename in expected_content: count += 1 assert not count @@ -194,59 +192,60 @@ def test_close_tmp_files(path): @pytest.mark.parametrize( "env_value, expected_parent_fn", [ - ("use_working_dir", os.getcwd), - (os.path.expanduser("~"), lambda: os.path.expanduser("~")), - (None, tempfile.gettempdir), + ("use_working_dir", lambda: Path.cwd()), + (Path("~").expanduser(), lambda: Path("~").expanduser()), + (None, lambda: Path(tempfile.gettempdir())), ], ) def test_get_trx_tmp_dir(env_value, expected_parent_fn, monkeypatch): if env_value is None: monkeypatch.delenv("TRX_TMPDIR", raising=False) else: - monkeypatch.setenv("TRX_TMPDIR", env_value) + monkeypatch.setenv("TRX_TMPDIR", str(env_value)) td = get_trx_tmp_dir() + tmp_path = Path(td.name) try: - assert os.path.dirname(td.name) == expected_parent_fn() - assert os.path.isdir(td.name) + assert tmp_path.parent == expected_parent_fn() + assert tmp_path.is_dir() finally: td.cleanup() - assert not os.path.isdir(td.name) + assert not tmp_path.is_dir() @pytest.mark.parametrize( "trx_tmpdir_env, expected_parent", [ - ("use_working_dir", lambda: os.getcwd()), - (os.path.expanduser("~"), lambda: os.path.expanduser("~")), - (None, lambda: tempfile.gettempdir()), + ("use_working_dir", lambda: Path.cwd()), + (Path("~").expanduser(), lambda: Path("~").expanduser()), + (None, lambda: Path(tempfile.gettempdir())), ], ) def test_change_tmp_dir(trx_tmpdir_env, expected_parent, monkeypatch): """Integration test through tmm.load(path), assuming that it eventually calls get_trx_tmp_dir().""" - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, "gs.trx") + gs_dir = get_home() / "gold_standard" + path = gs_dir / "gs.trx" if trx_tmpdir_env is None: monkeypatch.delenv("TRX_TMPDIR", raising=False) else: - monkeypatch.setenv("TRX_TMPDIR", trx_tmpdir_env) + monkeypatch.setenv("TRX_TMPDIR", str(trx_tmpdir_env)) tgm = tmm.load(path) tmp_gs_dir = deepcopy(tgm._uncompressed_folder_handle.name) - assert os.path.dirname(tmp_gs_dir) == expected_parent() + assert tmp_gs_dir.parent == expected_parent() tgm.close() - assert not os.path.isdir(tmp_gs_dir) + assert not tmp_gs_dir.is_dir() @pytest.mark.parametrize("path", ["gs.trx", "gs_fldr.trx"]) def test_complete_dir_from_trx(path): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, path) + gs_dir = get_home() / "gold_standard" + path = gs_dir / path tgm = tmm.load(path) if tgm._uncompressed_folder_handle is None: @@ -257,8 +256,8 @@ def test_complete_dir_from_trx(path): file_paths = [] for dirpath, _, filenames in os.walk(dir_to_check): for filename in filenames: - full_path = os.path.join(dirpath, filename) - cut_path = full_path.split(dir_to_check)[1][1:].replace("\\", "/") + full_path = Path(dirpath) / filename + cut_path = full_path.relative_to(dir_to_check).as_posix() file_paths.append(cut_path) expected_content = [ @@ -274,8 +273,8 @@ def test_complete_dir_from_trx(path): def test_complete_zip_from_trx(): - gs_dir = os.path.join(get_home(), "gold_standard") - path = os.path.join(gs_dir, "gs.trx") + gs_dir = get_home() / "gold_standard" + path = gs_dir / "gs.trx" with zipfile.ZipFile(path, mode="r") as zf: zip_file_list = zf.namelist() diff --git a/trx/tests/test_memmap.py b/trx/tests/test_memmap.py index 52dd05d7..1b4e4b8e 100644 --- a/trx/tests/test_memmap.py +++ b/trx/tests/test_memmap.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- import json -import os +from pathlib import Path import struct import zipfile @@ -127,7 +127,7 @@ def test_dichotomic_search(arr, l_bound, r_bound, expected): def test_create_memmap(basename, create, expected): if create: with get_trx_tmp_dir() as dirname: - filename = os.path.join(dirname, basename) + filename = Path(dirname) / basename fp = tmm._create_memmap( filename=filename, mode="w+", shape=(3, 4), dtype=np.int16 ) @@ -137,9 +137,9 @@ def test_create_memmap(basename, create, expected): else: with get_trx_tmp_dir() as dirname: - filename = os.path.join(dirname, basename) + filename = Path(dirname) / basename mmarr = tmm._create_memmap(filename=filename, shape=(0,), dtype=np.int16) - assert os.path.isfile(filename) + assert filename.is_file() assert np.array_equal(mmarr, np.zeros(shape=(0,), dtype=np.float32)) @@ -154,7 +154,7 @@ def test_create_memmap(basename, create, expected): ], ) def test_load(path, check_dpg, value_error): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path # Need to perhaps improve test if value_error: with pytest.raises(ValueError): @@ -167,19 +167,19 @@ def test_load(path, check_dpg, value_error): @pytest.mark.parametrize("path", ["small.trx"]) def test_load_zip(path): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path assert isinstance(tmm.load_from_zip(path), tmm.TrxFile) @pytest.mark.parametrize("path", ["small_fldr.trx"]) def test_load_directory(path): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path assert isinstance(tmm.load_from_directory(path), tmm.TrxFile) @pytest.mark.parametrize("path", ["small.trx"]) def test_concatenate(path): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path trx1 = tmm.load(path) trx2 = tmm.load(path) concat = tmm.concatenate([trx1, trx2]) @@ -192,7 +192,7 @@ def test_concatenate(path): @pytest.mark.parametrize("path", ["small.trx"]) def test_resize(path): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path trx1 = tmm.load(path) concat = tmm.TrxFile(nb_vertices=1000000, nb_streamlines=10000, init_as=trx1) @@ -206,7 +206,7 @@ def test_resize(path): @pytest.mark.parametrize("path, buffer", [("small.trx", 10000), ("small.trx", 0)]) def test_append(path, buffer): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path trx1 = tmm.load(path) concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=trx1) @@ -222,7 +222,7 @@ def test_append(path, buffer): @pytest.mark.parametrize("path, buffer", [("small.trx", 10000)]) @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed") def test_append_StatefulTractogram(path, buffer): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path tgm = tmm.load(path) obj = tgm.to_sft() concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=tgm) @@ -238,7 +238,7 @@ def test_append_StatefulTractogram(path, buffer): @pytest.mark.parametrize("path, buffer", [("small.trx", 10000)]) def test_append_Tractogram(path, buffer): - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path tgm = tmm.load(path) obj = tgm.to_tractogram() concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=tgm) @@ -300,7 +300,7 @@ def streamlines_func(): "dpv": {"fa": np.float16}, "dps": {"commit_weights": np.float32, "clusters_QB": np.uint16}, } - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path tgm = tmm.TrxFile.from_lazy_tractogram( obj, reference=path, extra_buffer=buffer, chunk_size=1000, dtype_dict=dtype_dict ) @@ -349,11 +349,11 @@ def test_initialize_empty_trx(tmp_path): assert tgm.header["NB_VERTICES"] == 0 assert len(tgm.streamlines) == 0 - out_path = os.path.join(tmp_path, "empty.trx") + out_path = tmp_path / "empty.trx" tmm.save(tgm, out_path) - assert os.path.exists(out_path) - file_size = os.path.getsize(out_path) + assert out_path.exists() + file_size = out_path.stat().st_size assert file_size < 500 # Should be very small, just header.json in zip with zipfile.ZipFile(out_path, "r") as zf: @@ -384,7 +384,7 @@ def test_trxfile_getgroup(): def test_trxfile_select(): - path = os.path.join(get_home(), "memmap_test_data", "small.trx") + path = get_home() / "memmap_test_data" / "small.trx" tgm = tmm.load(path) assert len(tgm.select([]).streamlines) == 0 @@ -399,10 +399,10 @@ def test_trxfile_select(): def test_save_after_select(tmp_path): - path = os.path.join(get_home(), "memmap_test_data", "small.trx") + path = get_home() / "memmap_test_data" / "small.trx" tgm = tmm.load(path) sub = tgm.select(list(range(5))) - out = os.path.join(tmp_path, "sub.trx") + out = tmp_path / "sub.trx" tmm.save(sub, out) loaded = tmm.load(out) assert len(loaded.streamlines) == 5 @@ -422,7 +422,7 @@ def test_trxfile_close(): @pytest.mark.parametrize("path", ["small.trx"]) def test_close_releases_mmap_from_zip(path): """close() must release mmap handles even when loaded via load_from_zip().""" - path = os.path.join(get_home(), "memmap_test_data", path) + path = get_home() / "memmap_test_data" / path tgm = tmm.load_from_zip(path) assert tgm._uncompressed_folder_handle is None @@ -542,7 +542,7 @@ def make_zip64_extra(orig_size, comp_size): _data = struct.pack(" np.dtype: """Convert a dtype to its little-endian equivalent. @@ -129,14 +136,14 @@ def is_sorted(a): return np.append(nib_offsets, nb_vertices).astype(nib_offsets.dtype) -def _generate_filename_from_data(arr: np.ndarray, filename: str) -> str: +def _generate_filename_from_data(arr: np.ndarray, filename: str | Path) -> str: """Determine the data type from array data and generate the appropriate filename. Parameters ---------- arr : np.ndarray A NumPy array (1-2D, otherwise ValueError raised). - filename : str + filename : str or Path The original filename. Returns @@ -144,7 +151,10 @@ def _generate_filename_from_data(arr: np.ndarray, filename: str) -> str: str An updated filename with appropriate extension. """ - base, ext = os.path.splitext(filename) + filename = Path(filename) + ext = "".join(filename.suffixes) + base = filename.name.removesuffix(ext) + ext = ext.lstrip(".") if ext: logging.warning("Will overwrite provided extension if needed.") @@ -165,12 +175,30 @@ def _generate_filename_from_data(arr: np.ndarray, filename: str) -> str: return new_filename -def _split_ext_with_dimensionality(filename: str) -> Tuple[str, int, str]: +def _get_last_ext(filename: Path) -> str: + """Return the last dot-segment as the extension. + + Handles dimensionality-encoded filenames like 'positions.3.float32'. + + Parameters + ---------- + filename : Path + Input filename. + + Returns + ------- + str + Last extension. + """ + return f".{filename.name.split('.')[-1]}" + + +def _split_ext_with_dimensionality(filename: str | Path) -> Tuple[str, int, str]: """Take a filename and split it into its components. Parameters ---------- - filename : str + filename : str or Path Input filename. Returns @@ -178,7 +206,7 @@ def _split_ext_with_dimensionality(filename: str) -> Tuple[str, int, str]: tuple A tuple of (basename, dimension, extension). """ - basename = os.path.basename(filename) + basename = Path(filename).name split = basename.split(".") if len(split) != 2 and len(split) != 3: @@ -274,7 +302,7 @@ def _dichotomic_search( def _create_memmap( - filename: str, + filename: str | Path, mode: str = "r", shape: Tuple = (1,), dtype: np.dtype = np.float32, @@ -285,7 +313,7 @@ def _create_memmap( Parameters ---------- - filename : str + filename : str or Path Filename where the empty memmap should be created. mode : str, optional File open mode (see np.memmap for options). Default is 'r'. @@ -303,8 +331,10 @@ def _create_memmap( np.ndarray Memory-mapped array or a zero-filled array if shape[0] is 0. """ + filename = Path(filename) + if np.dtype(dtype) == bool: - filename = filename.replace(".bool", ".bit") + filename = Path(str(filename).replace(".bool", ".bit")) # TRX format uses little-endian byte order for cross-platform compatibility dtype = _get_dtype_little_endian(dtype) @@ -314,18 +344,18 @@ def _create_memmap( filename, mode=mode, offset=offset, shape=shape, dtype=dtype, order=order ) else: - if not os.path.isfile(filename): + if not filename.is_file(): f = open(filename, "wb") f.close() return np.zeros(shape, dtype=dtype) -def load(input_obj: str, check_dpg: bool = True) -> Type["TrxFile"]: +def load(input_obj: str | Path, check_dpg: bool = True) -> Type["TrxFile"]: """Load a TrxFile (compressed or not). Parameters ---------- - input_obj : str + input_obj : str | Path A directory name or filepath to the TRX data. check_dpg : bool, optional Whether to check group metadata. Default is True. @@ -338,7 +368,8 @@ def load(input_obj: str, check_dpg: bool = True) -> Type["TrxFile"]: # TODO Check if 0 streamlines, then 0 vertices is expected (vice-versa) # TODO 4x4 affine matrices should contains values (no all-zeros) # TODO 3x1 dimensions array should contains values at each position (int) - if os.path.isfile(input_obj): + input_obj = Path(input_obj) + if input_obj.is_file(): was_compressed = False with zipfile.ZipFile(input_obj, "r") as zf: for info in zf.infolist(): @@ -356,7 +387,7 @@ def load(input_obj: str, check_dpg: bool = True) -> Type["TrxFile"]: ) else: tgm = load_from_zip(input_obj) - elif os.path.isdir(input_obj): + elif input_obj.is_dir(): tgm = load_from_directory(input_obj) else: raise ValueError("File/Folder does not exist") @@ -369,14 +400,14 @@ def load(input_obj: str, check_dpg: bool = True) -> Type["TrxFile"]: return tgm -def load_from_zip(filename: str) -> Type["TrxFile"]: +def load_from_zip(filename: str | Path) -> Type["TrxFile"]: """Load a TrxFile from a single zipfile. Note: Does not work with compressed zipfiles. Parameters ---------- - filename : str + filename : str or Path Path of the zipped TrxFile. Returns @@ -395,7 +426,7 @@ def load_from_zip(filename: str) -> Type["TrxFile"]: files_pointer_size = {} for zip_info in zf.filelist: elem_filename = zip_info.filename - _, ext = os.path.splitext(elem_filename) + ext = _get_last_ext(Path(elem_filename)) if ext == ".json" or zip_info.is_dir(): continue @@ -440,12 +471,12 @@ def load_from_zip(filename: str) -> Type["TrxFile"]: ) -def load_from_directory(directory: str) -> Type["TrxFile"]: +def load_from_directory(directory: str | Path) -> Type["TrxFile"]: """Load a TrxFile from a folder containing memmaps. Parameters ---------- - directory : str + directory : str or Path Path of the directory containing TRX data. Returns @@ -454,8 +485,8 @@ def load_from_directory(directory: str) -> Type["TrxFile"]: TrxFile representing the read data. """ - directory = os.path.abspath(directory) - with open(os.path.join(directory, "header.json")) as header: + directory = Path(directory).resolve() + with open(directory / "header.json") as header: header = json.load(header) header["VOXEL_TO_RASMM"] = np.reshape(header["VOXEL_TO_RASMM"], (4, 4)).astype( np.float32 @@ -463,9 +494,10 @@ def load_from_directory(directory: str) -> Type["TrxFile"]: header["DIMENSIONS"] = np.array(header["DIMENSIONS"], dtype=np.uint16) files_pointer_size = {} for root, _dirs, files in os.walk(directory): + root = Path(root) for name in files: - elem_filename = os.path.join(root, name) - _, ext = os.path.splitext(elem_filename) + elem_filename = root / name + ext = _get_last_ext(elem_filename) if ext == ".json": continue @@ -476,10 +508,10 @@ def load_from_directory(directory: str) -> Type["TrxFile"]: ext = ".bool" dtype_size = np.dtype(ext[1:]).itemsize - size = os.path.getsize(elem_filename) / dtype_size + size = elem_filename.stat().st_size / dtype_size if size.is_integer(): files_pointer_size[elem_filename] = 0, int(size) - elif os.path.getsize(elem_filename) == 1: + elif elem_filename.stat().st_size == 1: files_pointer_size[elem_filename] = 0, 0 else: raise ValueError("Wrong size or datatype") @@ -736,11 +768,11 @@ def _setup_groups_for_concatenation( tmp_dir = new_trx._uncompressed_folder_handle.name for group_key in all_groups_len.keys(): - if not os.path.isdir(os.path.join(tmp_dir, "groups/")): - os.mkdir(os.path.join(tmp_dir, "groups/")) + if not (Path(tmp_dir) / "groups").is_dir(): + os.mkdir(Path(tmp_dir) / "groups") dtype = all_groups_dtype[group_key] - group_filename = os.path.join(tmp_dir, f"groups/{group_key}.{dtype.name}") + group_filename = Path(tmp_dir) / f"groups/{group_key}.{dtype.name}" group_len = all_groups_len[group_key] new_trx.groups[group_key] = _create_memmap( group_filename, mode="w+", shape=(group_len,), dtype=dtype @@ -841,13 +873,13 @@ def save( ---------- tgm : TrxFile The TrxFile to save. - filename : str + filename : str or Path The path to save the TrxFile to. compression_standard : int, optional The compression standard to use, as defined by the ZipFile library. Default is zipfile.ZIP_STORED. """ - _, ext = os.path.splitext(filename) + ext = "".join(Path(filename).suffixes) if ext.lower() not in [".zip", ".trx", ""]: raise ValueError("Unsupported extension.") @@ -857,14 +889,16 @@ def save( if ext.lower() in [".zip", ".trx"]: zip_from_folder(tmp_dir_name, filename, compression_standard) else: - if os.path.isdir(filename): + if filename.is_dir(): shutil.rmtree(filename) shutil.copytree(tmp_dir_name, filename) copy_trx.close() def zip_from_folder( - directory: str, filename: str, compression_standard: Any = zipfile.ZIP_STORED + directory: str | Path, + filename: str | Path, + compression_standard: Any = zipfile.ZIP_STORED, ) -> None: """Zip on-disk memmaps into a single file. @@ -878,11 +912,14 @@ def zip_from_folder( The compression standard to use, as defined by the ZipFile library. Default is zipfile.ZIP_STORED. """ + directory = Path(directory) + filename = Path(filename) with zipfile.ZipFile(filename, mode="w", compression=compression_standard) as zf: for root, _, files in os.walk(directory): + root = Path(root) for name in files: - curr_filename = os.path.join(root, name) - tmp_filename = curr_filename.replace(directory, "")[1:] + curr_filename = root / name + tmp_filename = curr_filename.relative_to(directory) zf.write(curr_filename, tmp_filename) @@ -916,6 +953,7 @@ def __init__( init_as: Optional[Type["TrxFile"]] = None, reference: Union[ str, + Path, dict, Type[Nifti1Image], Type[TrkFile], @@ -1081,7 +1119,8 @@ def deepcopy(self) -> Type["TrxFile"]: # noqa: C901 A deepcopied TrxFile of the current TrxFile. """ tmp_dir = get_trx_tmp_dir() - out_json = open(os.path.join(tmp_dir.name, "header.json"), "w") + tmp_path = Path(tmp_dir.name) + out_json = open(tmp_path / "header.json", "w") tmp_header = deepcopy(self.header) if not isinstance(tmp_header["VOXEL_TO_RASMM"], list): @@ -1103,7 +1142,7 @@ def deepcopy(self) -> Type["TrxFile"]: # noqa: C901 # Only write positions and offsets if TRX is not empty if tmp_header["NB_STREAMLINES"] > 0 and tmp_header["NB_VERTICES"] > 0: positions_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "positions") + to_dump, tmp_path / "positions" ) _ensure_little_endian(to_dump).tofile(positions_filename) @@ -1116,12 +1155,12 @@ def deepcopy(self) -> Type["TrxFile"]: # noqa: C901 self.streamlines._offsets, self.header["NB_VERTICES"] ) offsets_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "offsets") + to_dump, tmp_path / "offsets" ) _ensure_little_endian(to_dump).tofile(offsets_filename) if len(self.data_per_vertex.keys()) > 0: - os.mkdir(os.path.join(tmp_dir.name, "dpv/")) + os.mkdir(tmp_path / "dpv") for dpv_key in self.data_per_vertex.keys(): if not self._copy_safe: to_dump = self.data_per_vertex[dpv_key].copy()._data @@ -1129,25 +1168,25 @@ def deepcopy(self) -> Type["TrxFile"]: # noqa: C901 to_dump = self.data_per_vertex[dpv_key]._data dpv_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "dpv/", dpv_key) + to_dump, tmp_path / "dpv" / dpv_key ) _ensure_little_endian(to_dump).tofile(dpv_filename) if len(self.data_per_streamline.keys()) > 0: - os.mkdir(os.path.join(tmp_dir.name, "dps/")) + os.mkdir(tmp_path / "dps") for dps_key in self.data_per_streamline.keys(): to_dump = self.data_per_streamline[dps_key] dps_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "dps/", dps_key) + to_dump, tmp_path / "dps" / dps_key ) _ensure_little_endian(to_dump).tofile(dps_filename) if len(self.groups.keys()) > 0: - os.mkdir(os.path.join(tmp_dir.name, "groups/")) + os.mkdir(tmp_path / "groups") for group_key in self.groups.keys(): to_dump = self.groups[group_key] group_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "groups/", group_key) + to_dump, tmp_path / "groups" / group_key ) _ensure_little_endian(to_dump).tofile(group_filename) @@ -1155,17 +1194,17 @@ def deepcopy(self) -> Type["TrxFile"]: # noqa: C901 continue for dpg_key in self.data_per_group[group_key].keys(): # Creates 'dpg/' only if required - if not os.path.isdir(os.path.join(tmp_dir.name, "dpg/")): - os.mkdir(os.path.join(tmp_dir.name, "dpg/")) - if not os.path.isdir(os.path.join(tmp_dir.name, "dpg/", group_key)): - os.mkdir(os.path.join(tmp_dir.name, "dpg/", group_key)) + if not (tmp_path / "dpg/").is_dir(): + os.mkdir(tmp_path / "dpg") + if not (tmp_path / "dpg" / group_key).is_dir(): + os.mkdir(tmp_path / "dpg" / group_key) to_dump = self.data_per_group[group_key][dpg_key] dpg_filename = _generate_filename_from_data( - to_dump, os.path.join(tmp_dir.name, "dpg/", group_key, dpg_key) + to_dump, tmp_path / "dpg" / group_key / dpg_key ) _ensure_little_endian(to_dump).tofile(dpg_filename) - copy_trx = load_from_directory(tmp_dir.name) + copy_trx = load_from_directory(tmp_path) copy_trx._uncompressed_folder_handle = tmp_dir return copy_trx @@ -1281,7 +1320,8 @@ def _initialize_empty_trx( # noqa: C901 """ tgm = TrxFile() tmp_dir = get_trx_tmp_dir() - logging.info(f"Temporary folder for memmaps: {tmp_dir.name}") + tmp_path = Path(tmp_dir.name) + logging.info(f"Temporary folder for memmaps: {tmp_path}") tgm.header["NB_VERTICES"] = nb_vertices tgm.header["NB_STREAMLINES"] = nb_streamlines @@ -1302,14 +1342,12 @@ def _initialize_empty_trx( # noqa: C901 logging.debug(f"Initializing lengths with dtype: {lengths_dtype.name}") # A TrxFile without init_as only contain the essential arrays - positions_filename = os.path.join( - tmp_dir.name, f"positions.3.{positions_dtype.name}" - ) + positions_filename = tmp_path / f"positions.3.{positions_dtype.name}" tgm.streamlines._data = _create_memmap( positions_filename, mode="w+", shape=(nb_vertices, 3), dtype=positions_dtype ) - offsets_filename = os.path.join(tmp_dir.name, f"offsets.{offsets_dtype.name}") + offsets_filename = tmp_path / f"offsets.{offsets_dtype.name}" tgm.streamlines._offsets = _create_memmap( offsets_filename, mode="w+", shape=(nb_streamlines,), dtype=offsets_dtype ) @@ -1320,24 +1358,20 @@ def _initialize_empty_trx( # noqa: C901 # Only the structure of fixed-size arrays is copied if init_as is not None: if len(init_as.data_per_vertex.keys()) > 0: - os.mkdir(os.path.join(tmp_dir.name, "dpv/")) + os.mkdir(tmp_path / "dpv") if len(init_as.data_per_streamline.keys()) > 0: - os.mkdir(os.path.join(tmp_dir.name, "dps/")) + os.mkdir(tmp_path / "dps") for dpv_key in init_as.data_per_vertex.keys(): dtype = init_as.data_per_vertex[dpv_key]._data.dtype tmp_as = init_as.data_per_vertex[dpv_key]._data if tmp_as.ndim == 1: - dpv_filename = os.path.join( - tmp_dir.name, f"dpv/{dpv_key}.{dtype.name}" - ) + dpv_filename = tmp_path / f"dpv/{dpv_key}.{dtype.name}" shape = (nb_vertices, 1) elif tmp_as.ndim == 2: dim = tmp_as.shape[-1] shape = (nb_vertices, dim) - dpv_filename = os.path.join( - tmp_dir.name, f"dpv/{dpv_key}.{dim}.{dtype.name}" - ) + dpv_filename = tmp_path / f"dpv/{dpv_key}.{dim}.{dtype.name}" else: raise ValueError("Invalid dimensionality.") @@ -1353,16 +1387,12 @@ def _initialize_empty_trx( # noqa: C901 dtype = init_as.data_per_streamline[dps_key].dtype tmp_as = init_as.data_per_streamline[dps_key] if tmp_as.ndim == 1: - dps_filename = os.path.join( - tmp_dir.name, f"dps/{dps_key}.{dtype.name}" - ) + dps_filename = tmp_path / f"dps/{dps_key}.{dtype.name}" shape = (nb_streamlines,) elif tmp_as.ndim == 2: dim = tmp_as.shape[-1] shape = (nb_streamlines, dim) - dps_filename = os.path.join( - tmp_dir.name, f"dps/{dps_key}.{dim}.{dtype.name}" - ) + dps_filename = tmp_path / f"dps/{dps_key}.{dim}.{dtype.name}" else: raise ValueError("Invalid dimensionality.") @@ -1380,8 +1410,8 @@ def _initialize_empty_trx( # noqa: C901 def _create_trx_from_pointer( # noqa: C901 header: dict, dict_pointer_size: dict, - root_zip: Optional[str] = None, - root: Optional[str] = None, + root_zip: str | Path | None = None, + root: str | Path | None = None, ) -> Type["TrxFile"]: """Create a TrxFile after reading the structure of a zip/folder. @@ -1392,9 +1422,9 @@ def _create_trx_from_pointer( # noqa: C901 dict_pointer_size : dict A dictionary containing the filenames of all the files within the TrxFile disk file/folder. - root_zip : str, optional + root_zip : str or Path, optional The path of the ZipFile pointer. - root : str, optional + root : str or Path, optional The dirname of the ZipFile pointer. Returns @@ -1409,6 +1439,7 @@ def _create_trx_from_pointer( # noqa: C901 if header["NB_STREAMLINES"] == 0 or header["NB_VERTICES"] == 0: return tgm + root = Path(root) if root is not None else None positions, offsets = None, None for elem_filename in dict_pointer_size.keys(): if root_zip: @@ -1416,30 +1447,33 @@ def _create_trx_from_pointer( # noqa: C901 else: filename = elem_filename - folder = os.path.dirname(elem_filename) + zip_path = PurePosixPath(elem_filename) + folder = zip_path.parent base, dim, ext = _split_ext_with_dimensionality(elem_filename) if ext == ".bit": ext = ".bool" mem_adress, size = dict_pointer_size[elem_filename] if root is not None: + disk_folder = Path(elem_filename).parent # This is for Unix - if os.name != "nt" and folder.startswith(root.rstrip("/")): - folder = folder.replace(root, "").lstrip("/") + if os.name != "nt" and disk_folder.is_relative_to(root): + relative_folder = disk_folder.relative_to(root) + folder = PurePosixPath(*relative_folder.parts) # These three are for Windows - elif os.path.isdir(folder) and os.path.basename(folder) in [ + elif disk_folder.is_dir() and disk_folder.name in [ "dpv", "dps", "groups", ]: - folder = os.path.basename(folder) - elif os.path.basename(os.path.dirname(folder)) == "dpg": - folder = os.path.join("dpg", os.path.basename(folder)) + folder = PurePosixPath(disk_folder.name) + elif disk_folder.parent.name == "dpg": + folder = DPG / folder.name else: - folder = "" + folder = ROOT # Parse/walk the directory tree - if base == "positions" and folder == "": + if base == "positions" and folder == ROOT: if size != tgm.header["NB_VERTICES"] * 3 or dim != 3: raise ValueError("Wrong data size/dimensionality.") positions = _create_memmap( @@ -1449,7 +1483,7 @@ def _create_trx_from_pointer( # noqa: C901 shape=(tgm.header["NB_VERTICES"], 3), dtype=ext[1:], ) - elif base == "offsets" and folder == "": + elif base == "offsets" and folder == ROOT: if size != tgm.header["NB_STREAMLINES"] + 1 or dim != 1: raise ValueError("Wrong offsets size/dimensionality.") offsets = _create_memmap( @@ -1463,7 +1497,7 @@ def _create_trx_from_pointer( # noqa: C901 lengths = _compute_lengths(offsets) else: lengths = [0] - elif folder == "dps": + elif folder == DPS: nb_scalar = size / tgm.header["NB_STREAMLINES"] if not nb_scalar.is_integer() or nb_scalar != dim: raise ValueError("Wrong dps size/dimensionality.") @@ -1473,7 +1507,7 @@ def _create_trx_from_pointer( # noqa: C901 tgm.data_per_streamline[base] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) - elif folder == "dpv": + elif folder == DPV: nb_scalar = size / tgm.header["NB_VERTICES"] if not nb_scalar.is_integer() or nb_scalar != dim: raise ValueError("Wrong dpv size/dimensionality.") @@ -1483,21 +1517,21 @@ def _create_trx_from_pointer( # noqa: C901 tgm.data_per_vertex[base] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) - elif folder.startswith("dpg"): + elif folder.parent == DPG: if int(size) != dim: raise ValueError("Wrong dpg size/dimensionality.") else: shape = (1, int(size)) # Handle the two-layers architecture - data_name = os.path.basename(base) - sub_folder = os.path.basename(folder) + data_name = base + sub_folder = folder.name if sub_folder not in tgm.data_per_group: tgm.data_per_group[sub_folder] = {} tgm.data_per_group[sub_folder][data_name] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) - elif folder == "groups": + elif folder == GROUPS: # Groups are simply indices, nothing else # TODO Crash if not uint? if dim != 1: @@ -1590,15 +1624,13 @@ def resize( # noqa: C901 else: tgm._copy_fixed_arrays_from(self) - tmp_dir = tgm._uncompressed_folder_handle.name + tmp_dir = Path(tgm._uncompressed_folder_handle.name) if len(self.groups.keys()) > 0: - os.mkdir(os.path.join(tmp_dir, "groups/")) + os.mkdir(tmp_dir / "groups/") for group_key in self.groups.keys(): group_dtype = self.groups[group_key].dtype - group_name = os.path.join( - tmp_dir, "groups/", f"{group_key}.{group_dtype.name}" - ) + group_name = tmp_dir / "groups" / f"{group_key}.{group_dtype.name}" ori_len = len(self.groups[group_key]) # Remove groups indices if resizing down @@ -1615,10 +1647,10 @@ def resize( # noqa: C901 return if len(self.data_per_group.keys()) > 0: - os.mkdir(os.path.join(tmp_dir, "dpg/")) + os.mkdir(tmp_dir / "dpg") for group_key in self.data_per_group: - if not os.path.isdir(os.path.join(tmp_dir, "dpg/", group_key)): - os.mkdir(os.path.join(tmp_dir, "dpg/", group_key)) + if not (tmp_dir / "dpg" / group_key).is_dir(): + os.mkdir(tmp_dir / "dpg" / group_key) if group_key not in tgm.data_per_group: tgm.data_per_group[group_key] = {} @@ -1626,7 +1658,7 @@ def resize( # noqa: C901 dpg_dtype = self.data_per_group[group_key][dpg_key].dtype dpg_filename = _generate_filename_from_data( self.data_per_group[group_key][dpg_key], - os.path.join(tmp_dir, "dpg/", group_key, dpg_key), + tmp_dir / "dpg" / group_key / dpg_key, ) shape = self.data_per_group[group_key][dpg_key].shape @@ -2015,9 +2047,10 @@ def from_sft(sft, dtype_dict=None): # For safety and for RAM, convert the whole object to memmaps tmp_dir = get_trx_tmp_dir() - save(tgm, tmp_dir.name) + tmp_path = Path(tmp_dir.name) + save(tgm, tmp_path) tgm.close() - tgm = load_from_directory(tmp_dir.name) + tgm = load_from_directory(tmp_path) tgm._uncompressed_folder_handle = tmp_dir sft.to_space(old_space) @@ -2112,10 +2145,11 @@ def from_tractogram( # For safety and for RAM, convert the whole object to memmaps tmp_dir = get_trx_tmp_dir() - save(tgm, tmp_dir.name) + tmp_path = Path(tmp_dir.name) + save(tgm, tmp_path) tgm.close() - tgm = load_from_directory(tmp_dir.name) + tgm = load_from_directory(tmp_path) del tmp_streamlines return tgm diff --git a/trx/utils.py b/trx/utils.py index 43608b41..fff95393 100644 --- a/trx/utils.py +++ b/trx/utils.py @@ -2,7 +2,7 @@ """Utility functions for reference handling, coordinate flips, and file operations.""" import logging -import os +from pathlib import Path import nibabel as nib from nibabel.streamlines.array_sequence import ArraySequence @@ -42,7 +42,7 @@ def split_name_with_gz(filename): Parameters ---------- - filename : str + filename : str or Path The filename to clean. Returns @@ -52,17 +52,19 @@ def split_name_with_gz(filename): ext : str The full extension. """ - base, ext = os.path.splitext(filename) + filename = Path(filename) + ext = filename.suffix + base = filename.with_suffix("") if ext == ".gz": # Test if we have a .nii additional extension - temp_base, add_ext = os.path.splitext(base) + add_ext = base.suffix if add_ext == ".nii" or add_ext == ".trk": ext = add_ext + ext - base = temp_base + base = base.with_suffix("") - return base, ext + return str(base), ext def get_reference_info_wrapper(reference): # noqa: C901 @@ -70,7 +72,7 @@ def get_reference_info_wrapper(reference): # noqa: C901 Parameters ---------- - reference : str or dict or Nifti1Image or TrkFile or Nifti1Header or TrxFile + reference : str or Path or dict or Nifti1Image or TrkFile or Nifti1Header or TrxFile Reference that provides the spatial attribute. Returns @@ -90,7 +92,8 @@ def get_reference_info_wrapper(reference): # noqa: C901 is_trk = False is_sft = False is_trx = False - if isinstance(reference, str): + if isinstance(reference, (str, Path)): + reference = Path(reference) _, ext = split_name_with_gz(reference) if ext in [".nii", ".nii.gz"]: header = nib.load(reference).header @@ -332,7 +335,7 @@ def load_matrix_in_any_format(filepath): Parameters ---------- - filepath : str + filepath : str or Path Path to the matrix file. Returns @@ -340,7 +343,7 @@ def load_matrix_in_any_format(filepath): matrix : numpy.ndarray The matrix. """ - _, ext = os.path.splitext(filepath) + ext = "".join(Path(filepath).suffixes) if ext == ".txt": data = np.loadtxt(filepath) elif ext == ".npy": diff --git a/trx/workflows.py b/trx/workflows.py index d5a23b77..8f41db20 100644 --- a/trx/workflows.py +++ b/trx/workflows.py @@ -7,6 +7,7 @@ import json import logging import os +from pathlib import Path import tempfile from nibabel.streamlines.array_sequence import ArraySequence @@ -276,7 +277,7 @@ def verify_header_compatibility(in_files): all_valid = True for filepath in in_files: - if not os.path.isfile(filepath): + if not filepath.is_file(): print(f"{filepath} does not exist") _, in_extension = split_name_with_gz(filepath) if in_extension not in [".trk", ".nii", ".nii.gz", ".trx"]: @@ -500,7 +501,7 @@ def _write_header(tmp_dir_name, reference, streamlines): if header["NB_STREAMLINES"] <= 1: raise IOError("To use this script, you need at least 2streamlines.") - with open(os.path.join(tmp_dir_name, "header.json"), "w") as out_json: + with open(Path(tmp_dir_name) / "header.json", "w") as out_json: json.dump(header, out_json) @@ -518,11 +519,11 @@ def _write_streamline_data(tmp_dir_name, streamlines, positions_dtype, offsets_d offsets_dtype : str Datatype for offsets array. """ - curr_filename = os.path.join(tmp_dir_name, f"positions.3.{positions_dtype}") + curr_filename = Path(tmp_dir_name) / f"positions.3.{positions_dtype}" positions = streamlines._data.astype(positions_dtype) tmm._ensure_little_endian(positions).tofile(curr_filename) - curr_filename = os.path.join(tmp_dir_name, f"offsets.{offsets_dtype}") + curr_filename = Path(tmp_dir_name) / f"offsets.{offsets_dtype}" offsets = streamlines._offsets.astype(offsets_dtype) tmm._ensure_little_endian(offsets).tofile(curr_filename) @@ -548,7 +549,7 @@ def _write_data_array(tmp_dir_name, subdir_name, args, is_dpg=False): Parameters ---------- - tmp_dir_name : str + tmp_dir_name : str | Path Base temporary directory. subdir_name : str Subdirectory name (dpv, dps, groups, dpg). @@ -562,16 +563,17 @@ def _write_data_array(tmp_dir_name, subdir_name, args, is_dpg=False): None Writes the array to disk. """ + tmp_dir_name = Path(tmp_dir_name) if is_dpg: - os.makedirs(os.path.join(tmp_dir_name, "dpg", args[0]), exist_ok=True) + os.makedirs(tmp_dir_name / "dpg" / args[0], exist_ok=True) curr_arr = load_matrix_in_any_format(args[1]).astype(args[2]) - basename = os.path.basename(os.path.splitext(args[1])[0]) + basename = Path(args[1]).stem dtype_str = _normalize_dtype(args[1]) if args[1] != "bool" else "bit" dtype = args[2] else: - os.makedirs(os.path.join(tmp_dir_name, subdir_name), exist_ok=True) + os.makedirs(tmp_dir_name / subdir_name, exist_ok=True) curr_arr = np.squeeze(load_matrix_in_any_format(args[0]).astype(args[1])) - basename = os.path.basename(os.path.splitext(args[0])[0]) + basename = Path(args[0]).stem dtype_str = _normalize_dtype(args[1]) dtype = dtype_str @@ -584,13 +586,9 @@ def _write_data_array(tmp_dir_name, subdir_name, args, is_dpg=False): dim = "" if curr_arr.ndim == 1 else f"{curr_arr.shape[-1]}." if is_dpg: - curr_filename = os.path.join( - tmp_dir_name, "dpg", args[0], f"{basename}.{dim}{dtype}" - ) + curr_filename = tmp_dir_name / "dpg" / args[0] / f"{basename}.{dim}{dtype}" else: - curr_filename = os.path.join( - tmp_dir_name, subdir_name, f"{basename}.{dim}{dtype}" - ) + curr_filename = tmp_dir_name / subdir_name / f"{basename}.{dim}{dtype}" tmm._ensure_little_endian(curr_arr).tofile(curr_filename)