From 494fff0bbc172cd7daabe7c329ec3c054f5452d6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jon=20Haitz=20Legarreta=20Gorro=C3=B1o?= Date: Sat, 22 Aug 2026 13:16:15 +0200 Subject: [PATCH] STY: Avoid using package name for variable or parameter names Avoid using package name for variable or parameter names: do not use `trx` as the name for parameters or variables. --- README.md | 4 +- docs/source/dev.md | 2 +- examples/plot_dps_dpv.py | 50 ++++---- examples/plot_groups.py | 48 ++++---- examples/plot_read_write_trx.py | 38 +++--- trx/cli.py | 10 +- trx/tests/test_cli.py | 86 ++++++------- trx/tests/test_io.py | 16 +-- trx/tests/test_memmap.py | 74 ++++++------ trx/trx_file_memmap.py | 206 ++++++++++++++++---------------- trx/workflows.py | 72 +++++------ 11 files changed, 303 insertions(+), 303 deletions(-) diff --git a/README.md b/README.md index b1a22d5e..42794415 100644 --- a/README.md +++ b/README.md @@ -33,10 +33,10 @@ pip install . from trx.io import load, save # Load a tractogram (supports .trx, .trk, .tck, .vtk, .fib, .dpy) -trx = load("tractogram.trx") +tgm = load("tractogram.trx") # Save to a different format -save(trx, "output.trk") +save(tgm, "output.trk") ``` ### Command-Line Interface diff --git a/docs/source/dev.md b/docs/source/dev.md index 19153466..5f92800c 100644 --- a/docs/source/dev.md +++ b/docs/source/dev.md @@ -251,7 +251,7 @@ def load(filename, reference=None): Examples -------- >>> from trx.io import load - >>> trx = load("tractogram.trx") + >>> tgm = load("tractogram.trx") """ pass ``` diff --git a/examples/plot_dps_dpv.py b/examples/plot_dps_dpv.py index 7457bf1e..f6a6bedf 100644 --- a/examples/plot_dps_dpv.py +++ b/examples/plot_dps_dpv.py @@ -52,10 +52,10 @@ trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") # Load the TRX file -trx = load(trx_path) +tgm = load(trx_path) -print(f"Loaded TRX with {len(trx)} streamlines") -print(f"Total vertices: {trx.header['NB_VERTICES']}") +print(f"Loaded TRX with {len(tgm)} streamlines") +print(f"Total vertices: {tgm.header['NB_VERTICES']}") # %% # Exploring Data Per Vertex (dpv) @@ -63,11 +63,11 @@ # # Let's see what dpv data is available. -print("Data Per Vertex keys:", list(trx.data_per_vertex.keys())) +print("Data Per Vertex keys:", list(tgm.data_per_vertex.keys())) # Examine each dpv field -for key in trx.data_per_vertex: - data = trx.data_per_vertex[key] +for key in tgm.data_per_vertex: + data = tgm.data_per_vertex[key] print(f"\n {key}:") print(f" Shape: {data._data.shape}") print(f" Dtype: {data._data.dtype}") @@ -80,9 +80,9 @@ # The dpv data is organized to match the streamlines. You can access # the dpv values for a specific streamline using the same indices. -if len(trx.data_per_vertex) > 0: - first_dpv_key = list(trx.data_per_vertex.keys())[0] - dpv_data = trx.data_per_vertex[first_dpv_key] +if len(tgm.data_per_vertex) > 0: + first_dpv_key = list(tgm.data_per_vertex.keys())[0] + dpv_data = tgm.data_per_vertex[first_dpv_key] # Get dpv values for the first streamline first_streamline_dpv = dpv_data[0] @@ -96,11 +96,11 @@ # # Now let's examine the dps data. -print("Data Per Streamline keys:", list(trx.data_per_streamline.keys())) +print("Data Per Streamline keys:", list(tgm.data_per_streamline.keys())) # Examine each dps field -for key in trx.data_per_streamline: - data = trx.data_per_streamline[key] +for key in tgm.data_per_streamline: + data = tgm.data_per_streamline[key] print(f"\n {key}:") print(f" Shape: {data.shape}") print(f" Dtype: {data.dtype}") @@ -113,10 +113,10 @@ # A common use case is filtering streamlines based on dps values. # For example, selecting streamlines with high FA values. -if len(trx.data_per_streamline) > 0: +if len(tgm.data_per_streamline) > 0: # Use the first dps key for demonstration - first_dps_key = list(trx.data_per_streamline.keys())[0] - dps_data = trx.data_per_streamline[first_dps_key] + first_dps_key = list(tgm.data_per_streamline.keys())[0] + dps_data = tgm.data_per_streamline[first_dps_key] # Calculate some statistics print(f"\nStatistics for '{first_dps_key}':") @@ -156,14 +156,14 @@ print("\nDemonstrating multi-dimensional data:") # Check for any multi-dimensional dpv -for key in trx.data_per_vertex: - data = trx.data_per_vertex[key] +for key in tgm.data_per_vertex: + data = tgm.data_per_vertex[key] if len(data._data.shape) > 1 and data._data.shape[1] > 1: print(f" {key}: {data._data.shape[1]}D data per vertex") # Check for any multi-dimensional dps -for key in trx.data_per_streamline: - data = trx.data_per_streamline[key] +for key in tgm.data_per_streamline: + data = tgm.data_per_streamline[key] if len(data.shape) > 1 and data.shape[1] > 1: print(f" {key}: {data.shape[1]}D data per streamline") @@ -177,14 +177,14 @@ # Get vertex counts for first few streamlines print("\nVertex distribution for first 5 streamlines:") -for i in range(min(5, len(trx))): - streamline = trx.streamlines[i] +for i in range(min(5, len(tgm))): + streamline = tgm.streamlines[i] print(f" Streamline {i}: {len(streamline)} vertices") # Total vertices should match -total_from_streamlines = sum(len(trx.streamlines[i]) for i in range(len(trx))) +total_from_streamlines = sum(len(tgm.streamlines[i]) for i in range(len(tgm))) print(f"\nTotal vertices from streamlines: {total_from_streamlines}") -print(f"Total vertices in header: {trx.header['NB_VERTICES']}") +print(f"Total vertices in header: {tgm.header['NB_VERTICES']}") # %% # Summary @@ -192,8 +192,8 @@ # # In this tutorial, you learned how to: # -# - Access dpv data using ``trx.data_per_vertex[key]`` -# - Access dps data using ``trx.data_per_streamline[key]`` +# - Access dpv data using ``tgm.data_per_vertex[key]`` +# - Access dps data using ``tgm.data_per_streamline[key]`` # - Understand the shape conventions for scalar and vector data # - Use metadata for statistical analysis # - Understand the file structure for dpv and dps diff --git a/examples/plot_groups.py b/examples/plot_groups.py index cddf871e..92eb9801 100644 --- a/examples/plot_groups.py +++ b/examples/plot_groups.py @@ -48,9 +48,9 @@ trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") # Load the TRX file -trx = load(trx_path) +tgm = load(trx_path) -print(f"Loaded TRX with {len(trx)} streamlines") +print(f"Loaded TRX with {len(tgm)} streamlines") # %% # Accessing groups @@ -59,15 +59,15 @@ # Groups are stored as a dictionary where keys are group names and values # are numpy arrays of streamline indices. -print(f"Available groups: {list(trx.groups.keys())}") +print(f"Available groups: {list(tgm.groups.keys())}") # Check the number of groups -print(f"Number of groups: {len(trx.groups)}") +print(f"Number of groups: {len(tgm.groups)}") # %% # Let's examine the groups in more detail: -for group_name, indices in trx.groups.items(): +for group_name, indices in tgm.groups.items(): print(f" {group_name}: {len(indices)} streamlines") # %% @@ -77,16 +77,16 @@ # You can extract all streamlines belonging to a specific group using # the ``get_group()`` method. -if len(trx.groups) > 0: +if len(tgm.groups) > 0: # Get the first group name - first_group = list(trx.groups.keys())[0] + first_group = list(tgm.groups.keys())[0] # Extract the group as a new TrxFile - group_trx = trx.get_group(first_group) - print(f"Extracted group '{first_group}' with {len(group_trx)} streamlines") + group_tgm = tgm.get_group(first_group) + print(f"Extracted group '{first_group}' with {len(group_tgm)} streamlines") # You can also access the raw indices - group_indices = trx.groups[first_group] + group_indices = tgm.groups[first_group] print(f"Raw indices (first 10): {group_indices[:10]}") else: print("No groups available in this file") @@ -98,12 +98,12 @@ # You can use group indices to select streamlines directly with the # ``select()`` method. -if len(trx.groups) > 0: - first_group = list(trx.groups.keys())[0] - indices = trx.groups[first_group] +if len(tgm.groups) > 0: + first_group = list(tgm.groups.keys())[0] + indices = tgm.groups[first_group] # Select streamlines using indices - selected = trx.select(indices[:5]) # Select first 5 from the group + selected = tgm.select(indices[:5]) # Select first 5 from the group print(f"Selected {len(selected)} streamlines from group '{first_group}'") # %% @@ -114,11 +114,11 @@ # This is useful for storing group-level statistics like mean FA, volume, # or color codes. -print(f"Data per group keys: {list(trx.data_per_group.keys())}") +print(f"Data per group keys: {list(tgm.data_per_group.keys())}") # Check what metadata is available for each group -for group_name in trx.data_per_group: - dpg_keys = list(trx.data_per_group[group_name].keys()) +for group_name in tgm.data_per_group: + dpg_keys = list(tgm.data_per_group[group_name].keys()) print(f" {group_name}: {dpg_keys}") # %% @@ -169,18 +169,18 @@ # A common workflow is to filter streamlines based on group membership # and then analyze or visualize specific bundles. -if len(trx.groups) > 0: +if len(tgm.groups) > 0: # Get all group names - group_names = list(trx.groups.keys()) + group_names = list(tgm.groups.keys()) # Report statistics for each group print("Group statistics:") for group_name in group_names: - group_trx = trx.get_group(group_name) - total_points = len(group_trx.streamlines._data) - avg_length = total_points / len(group_trx) if len(group_trx) > 0 else 0 + group_tgm = tgm.get_group(group_name) + total_points = len(group_tgm.streamlines._data) + avg_length = total_points / len(group_tgm) if len(group_tgm) > 0 else 0 print(f" {group_name}:") - print(f" - Streamlines: {len(group_trx)}") + print(f" - Streamlines: {len(group_tgm)}") print(f" - Total points: {total_points}") print(f" - Avg points per streamline: {avg_length:.1f}") @@ -190,7 +190,7 @@ # # In this tutorial, you learned how to: # -# - Access groups using ``trx.groups`` +# - Access groups using ``tgm.groups`` # - Extract group streamlines using ``get_group()`` # - Work with ``data_per_group`` (dpg) metadata # - Understand that groups can overlap diff --git a/examples/plot_read_write_trx.py b/examples/plot_read_write_trx.py index f9af07ff..a9fa10d7 100644 --- a/examples/plot_read_write_trx.py +++ b/examples/plot_read_write_trx.py @@ -36,7 +36,7 @@ trx_path = os.path.join(trx_home, "gold_standard", "gs.trx") # Load the TRX file -trx = load(trx_path) +tgm = load(trx_path) print("TRX file loaded successfully!") @@ -48,16 +48,16 @@ # Let's look at what's inside our loaded file. # Print a summary of the TRX file -print(trx) +print(tgm) # %% # The header contains essential metadata about the tractogram: print("Header information:") -print(f" Number of streamlines: {trx.header['NB_STREAMLINES']}") -print(f" Number of vertices: {trx.header['NB_VERTICES']}") -print(f" Image dimensions: {trx.header['DIMENSIONS']}") -print(f" Voxel to RASMM affine:\n{trx.header['VOXEL_TO_RASMM']}") +print(f" Number of streamlines: {tgm.header['NB_STREAMLINES']}") +print(f" Number of vertices: {tgm.header['NB_VERTICES']}") +print(f" Image dimensions: {tgm.header['DIMENSIONS']}") +print(f" Voxel to RASMM affine:\n{tgm.header['VOXEL_TO_RASMM']}") # %% # Accessing streamlines @@ -66,11 +66,11 @@ # Streamlines are the core data in a TRX file. Each streamline is a sequence # of 3D points representing a fiber tract in the brain. -print(f"Number of streamlines: {len(trx)}") -print(f"Total number of vertices: {len(trx.streamlines._data)}") +print(f"Number of streamlines: {len(tgm)}") +print(f"Total number of vertices: {len(tgm.streamlines._data)}") # Access the first streamline -first_streamline = trx.streamlines[0] +first_streamline = tgm.streamlines[0] print(f"\nFirst streamline has {len(first_streamline)} points") print(f"First 3 points of the first streamline:\n{first_streamline[:3]}") @@ -80,9 +80,9 @@ # # TRX files can contain additional data per vertex (dpv) and per streamline (dps). -print("Data per vertex (dpv) keys:", list(trx.data_per_vertex.keys())) -print("Data per streamline (dps) keys:", list(trx.data_per_streamline.keys())) -print("Groups:", list(trx.groups.keys())) +print("Data per vertex (dpv) keys:", list(tgm.data_per_vertex.keys())) +print("Data per streamline (dps) keys:", list(tgm.data_per_streamline.keys())) +print("Groups:", list(tgm.groups.keys())) # %% # Selecting a subset of streamlines @@ -91,13 +91,13 @@ # You can easily select a subset of streamlines using indices or slicing. # Select first 5 streamlines -subset = trx[:5] +subset = tgm[:5] print(f"Subset has {len(subset)} streamlines") # Select specific streamlines by indices (ensure indices are valid) -max_idx = len(trx) - 1 +max_idx = len(tgm) - 1 indices = [0, min(2, max_idx), min(5, max_idx)] -selected = trx.select(indices) +selected = tgm.select(indices) print(f"Selected {len(selected)} streamlines") # %% @@ -110,7 +110,7 @@ with tempfile.TemporaryDirectory() as tmpdir: # Save as TRX file (zip archive) output_path = os.path.join(tmpdir, "output.trx") - save(trx, output_path) + save(tgm, output_path) print(f"Saved TRX file to: {output_path}") print(f"File size: {os.path.getsize(output_path)} bytes") @@ -126,10 +126,10 @@ # preserving the spatial reference information. # Create a deepcopy of the loaded TRX file -trx_copy = trx.deepcopy() +tgm_copy = tgm.deepcopy() -print(f"Created copy with {len(trx_copy)} streamlines") -print(f"Header preserved: DIMENSIONS = {trx_copy.header['DIMENSIONS']}") +print(f"Created copy with {len(tgm_copy)} streamlines") +print(f"Header preserved: DIMENSIONS = {tgm_copy.header['DIMENSIONS']}") # %% # Summary diff --git a/trx/cli.py b/trx/cli.py index 079b17be..77a4d62e 100644 --- a/trx/cli.py +++ b/trx/cli.py @@ -233,7 +233,7 @@ def concatenate_tractograms( has_group = True trx_list.append(tractogram_obj) - trx = concatenate( + tgm = concatenate( trx_list, delete_dpv=delete_dpv, delete_dps=delete_dps, @@ -241,7 +241,7 @@ def concatenate_tractograms( check_space_attributes=True, preallocation=False, ) - save(trx, str(out_tractogram)) + save(tgm, str(out_tractogram)) typer.echo( typer.style( @@ -1002,8 +1002,8 @@ def info( typer.echo("") # Show TRX content info - trx = load_trx(str(in_tractogram)) - typer.echo(trx) + tgm = load_trx(str(in_tractogram)) + typer.echo(tgm) # Show file listing (unzip -l style) typer.echo("\nArchive contents:") @@ -1021,7 +1021,7 @@ def info( typer.echo("--------- -------") typer.echo(f"{total_uncompressed:>9} {num_files} files") - trx.close() + tgm.close() def main(): diff --git a/trx/tests/test_cli.py b/trx/tests/test_cli.py index 6c6fa28e..d8fecfda 100644 --- a/trx/tests/test_cli.py +++ b/trx/tests/test_cli.py @@ -46,7 +46,7 @@ def test_create_temp_memmap_uses_reopenable_path(tmp_path): def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): - trx = SimpleNamespace( + tgm = SimpleNamespace( streamlines=SimpleNamespace( _data=np.arange(6, dtype=np.float16).reshape((2, 3)), _offsets=np.array([0, 3], dtype=np.uint64), @@ -62,14 +62,14 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): }, groups={"mock_group": np.array([0, 1], dtype=np.int32)}, ) - trx.close = lambda: None + tgm.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.load", return_value=tgm), patch("trx.workflows.tmm.save") as mock_save, patch( "trx.workflows.tempfile.NamedTemporaryFile", @@ -91,13 +91,13 @@ def test_manipulate_trx_datatype_uses_reopenable_memmaps(tmp_path): }, ) - 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") + assert tgm.streamlines._data.dtype == np.dtype("float32") + assert tgm.streamlines._offsets.dtype == np.dtype("uint32") + assert tgm.data_per_vertex["mock_dpv"]._data.dtype == np.dtype("uint16") + assert tgm.data_per_streamline["mock_dps"].dtype == np.dtype("float32") + assert tgm.data_per_group["mock_group"]["mock_dpg"].dtype == np.dtype("float64") + assert tgm.groups["mock_group"].dtype == np.dtype("uint16") + mock_save.assert_called_once_with(tgm, "out.trx") def _normalize_dtype_dict(dtype_dict): @@ -265,12 +265,12 @@ def test_execution_convert_to_trx(self, tmp_path): data_fix = np.load(exp_data) offsets_fix = np.load(exp_offsets) - trx = tmm.load(out_trx_path) - assert_equal(trx.streamlines._data.dtype, np.float32) - assert_equal(trx.streamlines._offsets.dtype, np.uint32) - assert_array_equal(trx.streamlines._data, data_fix) - assert_array_equal(trx.streamlines._offsets, offsets_fix) - trx.close() + tgm = tmm.load(out_trx_path) + assert_equal(tgm.streamlines._data.dtype, np.float32) + assert_equal(tgm.streamlines._offsets.dtype, np.uint32) + assert_array_equal(tgm.streamlines._data, data_fix) + assert_array_equal(tgm.streamlines._offsets, offsets_fix) + tgm.close() @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_from_trx(self, tmp_path): @@ -310,10 +310,10 @@ def test_execution_convert_dtype_p16_o64(self, tmp_path): offsets_dtype="uint64", ) - trx = tmm.load(out_convert_path) - assert_equal(trx.streamlines._data.dtype, np.float16) - assert_equal(trx.streamlines._offsets.dtype, np.uint64) - trx.close() + tgm = tmm.load(out_convert_path) + assert_equal(tgm.streamlines._data.dtype, np.float16) + assert_equal(tgm.streamlines._offsets.dtype, np.uint64) + tgm.close() @pytest.mark.skipif(not dipy_available, reason="Dipy is not installed.") def test_execution_convert_dtype_p64_o32(self, tmp_path): @@ -327,10 +327,10 @@ def test_execution_convert_dtype_p64_o32(self, tmp_path): offsets_dtype="uint32", ) - trx = tmm.load(out_convert_path) - assert_equal(trx.streamlines._data.dtype, np.float64) - assert_equal(trx.streamlines._offsets.dtype, np.uint32) - trx.close() + tgm = tmm.load(out_convert_path) + assert_equal(tgm.streamlines._data.dtype, np.float64) + assert_equal(tgm.streamlines._offsets.dtype, np.uint32) + 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") @@ -421,42 +421,42 @@ def test_execution_generate_trx_from_scratch(self, tmp_path): 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")) - trx = tmm.concatenate([trx1, trx2], preallocation=False) + tgm = tmm.concatenate([trx1, trx2], preallocation=False) # Right size - assert_equal(len(trx.streamlines), 2 * len(trx1.streamlines)) + assert_equal(len(tgm.streamlines), 2 * len(trx1.streamlines)) # Right data end_idx = trx1.header["NB_VERTICES"] - assert_allclose(trx.streamlines._data[:end_idx], trx1.streamlines._data) - assert_allclose(trx.streamlines._data[end_idx:], trx2.streamlines._data) + assert_allclose(tgm.streamlines._data[:end_idx], trx1.streamlines._data) + assert_allclose(tgm.streamlines._data[end_idx:], trx2.streamlines._data) # Right data_per_* - for key in trx.data_per_vertex.keys(): + for key in tgm.data_per_vertex.keys(): assert_equal( - trx.data_per_vertex[key]._data[:end_idx], + tgm.data_per_vertex[key]._data[:end_idx], trx1.data_per_vertex[key]._data, ) assert_equal( - trx.data_per_vertex[key]._data[end_idx:], + tgm.data_per_vertex[key]._data[end_idx:], trx2.data_per_vertex[key]._data, ) end_idx = trx1.header["NB_STREAMLINES"] - for key in trx.data_per_streamline.keys(): + for key in tgm.data_per_streamline.keys(): assert_equal( - trx.data_per_streamline[key][:end_idx], + tgm.data_per_streamline[key][:end_idx], trx1.data_per_streamline[key], ) assert_equal( - trx.data_per_streamline[key][end_idx:], + tgm.data_per_streamline[key][end_idx:], trx2.data_per_streamline[key], ) # Validate out_concat_path = os.path.join(tmp_path, "concat.trx") out_valid_path = os.path.join(tmp_path, "valid.trx") - tmm.save(trx, out_concat_path) + tmm.save(tgm, out_concat_path) validate_tractogram( out_concat_path, None, @@ -467,10 +467,10 @@ def test_execution_concatenate_validate_trx(self, tmp_path): trx_val = tmm.load(out_valid_path) # Right dtype and size - assert DeepDiff(trx.get_dtype_dict(), trx_val.get_dtype_dict()) == {} + assert DeepDiff(tgm.get_dtype_dict(), trx_val.get_dtype_dict()) == {} assert_equal(len(trx1.streamlines), len(trx_val.streamlines)) - trx.close() + tgm.close() trx1.close() trx2.close() trx_val.close() @@ -478,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") - trx = tmm.load(expected_trx) + tgm = tmm.load(expected_trx) expected_dtype = { "positions": np.dtype("float16"), @@ -501,12 +501,12 @@ def test_execution_manipulate_trx_datatype(self, tmp_path): assert ( DeepDiff( - trx.get_dtype_dict(), + tgm.get_dtype_dict(), _normalize_dtype_dict(expected_dtype), ) == {} ) - trx.close() + tgm.close() generated_dtype = { "positions": np.dtype("float32"), @@ -535,12 +535,12 @@ def test_execution_manipulate_trx_datatype(self, tmp_path): ), ): manipulate_trx_datatype(expected_trx, out_gen_path, generated_dtype) - trx = tmm.load(out_gen_path) + tgm = tmm.load(out_gen_path) assert ( DeepDiff( - trx.get_dtype_dict(), + tgm.get_dtype_dict(), _normalize_dtype_dict(generated_dtype), ) == {} ) - trx.close() + tgm.close() diff --git a/trx/tests/test_io.py b/trx/tests/test_io.py index 75ec69d1..4b0678a4 100644 --- a/trx/tests/test_io.py +++ b/trx/tests/test_io.py @@ -159,7 +159,7 @@ def test_close_tmp_files(path): gs_dir = os.path.join(get_home(), "gold_standard") path = os.path.join(gs_dir, path) - trx = tmm.load(path) + tgm = tmm.load(path) process = psutil.Process(os.getpid()) open_files = process.open_files() @@ -180,7 +180,7 @@ def test_close_tmp_files(path): count += 1 assert count == 6 - trx.close() + tgm.close() open_files = process.open_files() count = 0 @@ -234,12 +234,12 @@ def test_change_tmp_dir(trx_tmpdir_env, expected_parent, monkeypatch): else: monkeypatch.setenv("TRX_TMPDIR", trx_tmpdir_env) - trx = tmm.load(path) - tmp_gs_dir = deepcopy(trx._uncompressed_folder_handle.name) + tgm = tmm.load(path) + tmp_gs_dir = deepcopy(tgm._uncompressed_folder_handle.name) assert os.path.dirname(tmp_gs_dir) == expected_parent() - trx.close() + tgm.close() assert not os.path.isdir(tmp_gs_dir) @@ -248,11 +248,11 @@ def test_complete_dir_from_trx(path): gs_dir = os.path.join(get_home(), "gold_standard") path = os.path.join(gs_dir, path) - trx = tmm.load(path) - if trx._uncompressed_folder_handle is None: + tgm = tmm.load(path) + if tgm._uncompressed_folder_handle is None: dir_to_check = path else: - dir_to_check = trx._uncompressed_folder_handle.name + dir_to_check = tgm._uncompressed_folder_handle.name file_paths = [] for dirpath, _, filenames in os.walk(dir_to_check): diff --git a/trx/tests/test_memmap.py b/trx/tests/test_memmap.py index 59eb298d..52dd05d7 100644 --- a/trx/tests/test_memmap.py +++ b/trx/tests/test_memmap.py @@ -223,32 +223,32 @@ def test_append(path, buffer): @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) - trx = tmm.load(path) - obj = trx.to_sft() - concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=trx) + tgm = tmm.load(path) + obj = tgm.to_sft() + concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=tgm) concat.append(obj, extra_buffer=buffer) if buffer > 0: concat.resize() assert len(concat) == len(obj) - trx.close() + tgm.close() concat.close() @pytest.mark.parametrize("path, buffer", [("small.trx", 10000)]) def test_append_Tractogram(path, buffer): path = os.path.join(get_home(), "memmap_test_data", path) - trx = tmm.load(path) - obj = trx.to_tractogram() - concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=trx) + tgm = tmm.load(path) + obj = tgm.to_tractogram() + concat = tmm.TrxFile(nb_vertices=1, nb_streamlines=1, init_as=tgm) concat.append(obj, extra_buffer=buffer) if buffer > 0: concat.resize() assert len(concat) == len(obj) - trx.close() + tgm.close() concat.close() @@ -301,11 +301,11 @@ def streamlines_func(): "dps": {"commit_weights": np.float32, "clusters_QB": np.uint16}, } path = os.path.join(get_home(), "memmap_test_data", path) - trx = tmm.TrxFile.from_lazy_tractogram( + tgm = tmm.TrxFile.from_lazy_tractogram( obj, reference=path, extra_buffer=buffer, chunk_size=1000, dtype_dict=dtype_dict ) - assert len(trx) == len(gen_range) + assert len(tgm) == len(gen_range) def test_zip_from_folder(): @@ -344,13 +344,13 @@ def test_copy_fixed_arrays_from(): def test_initialize_empty_trx(tmp_path): """Test creating, saving, and loading an empty TRX file.""" - trx = tmm.TrxFile() - assert trx.header["NB_STREAMLINES"] == 0 - assert trx.header["NB_VERTICES"] == 0 - assert len(trx.streamlines) == 0 + tgm = tmm.TrxFile() + assert tgm.header["NB_STREAMLINES"] == 0 + assert tgm.header["NB_VERTICES"] == 0 + assert len(tgm.streamlines) == 0 out_path = os.path.join(tmp_path, "empty.trx") - tmm.save(trx, out_path) + tmm.save(tgm, out_path) assert os.path.exists(out_path) file_size = os.path.getsize(out_path) @@ -385,30 +385,30 @@ def test_trxfile_getgroup(): def test_trxfile_select(): path = os.path.join(get_home(), "memmap_test_data", "small.trx") - trx = tmm.load(path) + tgm = tmm.load(path) - assert len(trx.select([]).streamlines) == 0 - assert len(trx.select([0]).streamlines) == 1 + assert len(tgm.select([]).streamlines) == 0 + assert len(tgm.select([0]).streamlines) == 1 idx = list(range(10)) - sub = trx.select(idx) + sub = tgm.select(idx) assert len(sub.streamlines) == len(idx) assert not sub._copy_safe - trx.close() + tgm.close() def test_save_after_select(tmp_path): path = os.path.join(get_home(), "memmap_test_data", "small.trx") - trx = tmm.load(path) - sub = trx.select(list(range(5))) + tgm = tmm.load(path) + sub = tgm.select(list(range(5))) out = os.path.join(tmp_path, "sub.trx") tmm.save(sub, out) loaded = tmm.load(out) assert len(loaded.streamlines) == 5 assert len(loaded.streamlines._data) == len(sub.streamlines.copy()._data) loaded.close() - trx.close() + tgm.close() def test_trxfile_to_memory(): @@ -423,15 +423,15 @@ def test_trxfile_close(): 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) - trx = tmm.load_from_zip(path) + tgm = tmm.load_from_zip(path) - assert trx._uncompressed_folder_handle is None + assert tgm._uncompressed_folder_handle is None - mmap_obj = trx.streamlines._data._mmap + mmap_obj = tgm.streamlines._data._mmap assert mmap_obj is not None, "expected a live mmap before close()" assert not mmap_obj.closed, "mmap should be open before close()" - trx.close() + tgm.close() assert mmap_obj.closed, ( "mmap is still open after close() — the mmap teardown was skipped " @@ -623,11 +623,11 @@ def make_zip64_extra(orig_size, comp_size): ) ) - trx = tmm.load_from_zip(trx_path) - np.testing.assert_array_almost_equal(trx.streamlines._data, positions) - assert trx.header["NB_VERTICES"] == 2 - assert trx.header["NB_STREAMLINES"] == 1 - trx.close() + tgm = tmm.load_from_zip(trx_path) + np.testing.assert_array_almost_equal(tgm.streamlines._data, positions) + assert tgm.header["NB_VERTICES"] == 2 + assert tgm.header["NB_STREAMLINES"] == 1 + tgm.close() def test_load_zip_with_local_header_extra_field(tmp_path): @@ -727,12 +727,12 @@ def test_load_zip_with_local_header_extra_field(tmp_path): ) ) - trx = tmm.load_from_zip(trx_path) - np.testing.assert_array_almost_equal(trx.streamlines._data, positions) - assert trx.header["NB_VERTICES"] == 2 - assert trx.header["NB_STREAMLINES"] == 1 + tgm = tmm.load_from_zip(trx_path) + np.testing.assert_array_almost_equal(tgm.streamlines._data, positions) + assert tgm.header["NB_VERTICES"] == 2 + assert tgm.header["NB_STREAMLINES"] == 1 - trx.close() + tgm.close() def test_endianness_roundtrip(): diff --git a/trx/trx_file_memmap.py b/trx/trx_file_memmap.py index 665ae7ba..5ce575e0 100644 --- a/trx/trx_file_memmap.py +++ b/trx/trx_file_memmap.py @@ -349,24 +349,24 @@ def load(input_obj: str, check_dpg: bool = True) -> Type["TrxFile"]: with zipfile.ZipFile(input_obj, "r") as zf: tmp_dir = get_trx_tmp_dir() zf.extractall(tmp_dir.name) - trx = load_from_directory(tmp_dir.name) - trx._uncompressed_folder_handle = tmp_dir + tgm = load_from_directory(tmp_dir.name) + tgm._uncompressed_folder_handle = tmp_dir logging.info( "File was compressed, call the close() function before exiting." ) else: - trx = load_from_zip(input_obj) + tgm = load_from_zip(input_obj) elif os.path.isdir(input_obj): - trx = load_from_directory(input_obj) + tgm = load_from_directory(input_obj) else: raise ValueError("File/Folder does not exist") # Example of robust check for metadata if check_dpg: - for dpg in trx.data_per_group.keys(): - if dpg not in trx.groups.keys(): + for dpg in tgm.data_per_group.keys(): + if dpg not in tgm.groups.keys(): raise ValueError(f"An undeclared group ({dpg}) has data_per_group.") - return trx + return tgm def load_from_zip(filename: str) -> Type["TrxFile"]: @@ -833,13 +833,13 @@ def concatenate( def save( - trx: "TrxFile", filename: str, compression_standard: Any = zipfile.ZIP_STORED + tgm: "TrxFile", filename: str, compression_standard: Any = zipfile.ZIP_STORED ) -> None: """Save a TrxFile (compressed or not). Parameters ---------- - trx : TrxFile + tgm : TrxFile The TrxFile to save. filename : str The path to save the TrxFile to. @@ -851,7 +851,7 @@ def save( if ext.lower() not in [".zip", ".trx", ""]: raise ValueError("Unsupported extension.") - copy_trx = trx.deepcopy() + copy_trx = tgm.deepcopy() copy_trx.resize() tmp_dir_name = copy_trx._uncompressed_folder_handle.name if ext.lower() in [".zip", ".trx"]: @@ -976,10 +976,10 @@ def __init__( f"Preallocating TrxFile with size {nb_streamlines} streamlines and " f"{nb_vertices} vertices." ) - trx = self._initialize_empty_trx( + tgm = self._initialize_empty_trx( nb_streamlines, nb_vertices, init_as=init_as ) - self.__dict__ = trx.__dict__ + self.__dict__ = tgm.__dict__ else: raise ValueError("You must declare both nb_vertices AND NB_STREAMLINES") @@ -1192,7 +1192,7 @@ def _get_real_len(self) -> Tuple[int, int]: def _copy_fixed_arrays_from( self, - trx: Type["TrxFile"], + tgm: Type["TrxFile"], strs_start: int = 0, pts_start: int = 0, nb_strs_to_copy: Optional[int] = None, @@ -1201,7 +1201,7 @@ def _copy_fixed_arrays_from( Parameters ---------- - trx : TrxFile + tgm : TrxFile TrxFile to copy data from. strs_start : int, optional The start index of the streamline. Default is 0. @@ -1217,10 +1217,10 @@ def _copy_fixed_arrays_from( streamlines and end of copied points. """ if nb_strs_to_copy is None: - curr_strs_len, curr_pts_len = trx._get_real_len() + curr_strs_len, curr_pts_len = tgm._get_real_len() else: curr_strs_len = int(nb_strs_to_copy) - curr_pts_len = np.sum(trx.streamlines._lengths[0:curr_strs_len]) + curr_pts_len = np.sum(tgm.streamlines._lengths[0:curr_strs_len]) curr_pts_len = int(curr_pts_len) strs_end = strs_start + curr_strs_len @@ -1230,27 +1230,27 @@ def _copy_fixed_arrays_from( return strs_start, pts_start # Mandatory arrays - self.streamlines._data[pts_start:pts_end] = trx.streamlines._data[ + self.streamlines._data[pts_start:pts_end] = tgm.streamlines._data[ 0:curr_pts_len ] self.streamlines._offsets[strs_start:strs_end] = ( - trx.streamlines._offsets[0:curr_strs_len] + pts_start + tgm.streamlines._offsets[0:curr_strs_len] + pts_start ) - self.streamlines._lengths[strs_start:strs_end] = trx.streamlines._lengths[ + self.streamlines._lengths[strs_start:strs_end] = tgm.streamlines._lengths[ 0:curr_strs_len ] # Optional fixed-sized arrays for dpv_key in self.data_per_vertex.keys(): self.data_per_vertex[dpv_key]._data[pts_start:pts_end] = ( - trx.data_per_vertex[dpv_key]._data[0:curr_pts_len] + tgm.data_per_vertex[dpv_key]._data[0:curr_pts_len] ) self.data_per_vertex[dpv_key]._offsets = self.streamlines._offsets self.data_per_vertex[dpv_key]._lengths = self.streamlines._lengths for dps_key in self.data_per_streamline.keys(): self.data_per_streamline[dps_key][strs_start:strs_end] = ( - trx.data_per_streamline[dps_key][0:curr_strs_len] + tgm.data_per_streamline[dps_key][0:curr_strs_len] ) return strs_end, pts_end @@ -1279,16 +1279,16 @@ def _initialize_empty_trx( # noqa: C901 TrxFile An empty TrxFile preallocated with a certain size. """ - trx = TrxFile() + tgm = TrxFile() tmp_dir = get_trx_tmp_dir() logging.info(f"Temporary folder for memmaps: {tmp_dir.name}") - trx.header["NB_VERTICES"] = nb_vertices - trx.header["NB_STREAMLINES"] = nb_streamlines + tgm.header["NB_VERTICES"] = nb_vertices + tgm.header["NB_STREAMLINES"] = nb_streamlines if init_as is not None: - trx.header["VOXEL_TO_RASMM"] = init_as.header["VOXEL_TO_RASMM"] - trx.header["DIMENSIONS"] = init_as.header["DIMENSIONS"] + tgm.header["VOXEL_TO_RASMM"] = init_as.header["VOXEL_TO_RASMM"] + tgm.header["DIMENSIONS"] = init_as.header["DIMENSIONS"] positions_dtype = init_as.streamlines._data.dtype offsets_dtype = init_as.streamlines._offsets.dtype lengths_dtype = init_as.streamlines._lengths.dtype @@ -1305,15 +1305,15 @@ def _initialize_empty_trx( # noqa: C901 positions_filename = os.path.join( tmp_dir.name, f"positions.3.{positions_dtype.name}" ) - trx.streamlines._data = _create_memmap( + 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}") - trx.streamlines._offsets = _create_memmap( + tgm.streamlines._offsets = _create_memmap( offsets_filename, mode="w+", shape=(nb_streamlines,), dtype=offsets_dtype ) - trx.streamlines._lengths = np.zeros( + tgm.streamlines._lengths = np.zeros( shape=(nb_streamlines,), dtype=lengths_dtype ) @@ -1342,12 +1342,12 @@ def _initialize_empty_trx( # noqa: C901 raise ValueError("Invalid dimensionality.") logging.debug(f"Initializing {dpv_key} (dpv) with dtype: {dtype.name}") - trx.data_per_vertex[dpv_key] = ArraySequence() - trx.data_per_vertex[dpv_key]._data = _create_memmap( + tgm.data_per_vertex[dpv_key] = ArraySequence() + tgm.data_per_vertex[dpv_key]._data = _create_memmap( dpv_filename, mode="w+", shape=shape, dtype=dtype ) - trx.data_per_vertex[dpv_key]._offsets = trx.streamlines._offsets - trx.data_per_vertex[dpv_key]._lengths = trx.streamlines._lengths + tgm.data_per_vertex[dpv_key]._offsets = tgm.streamlines._offsets + tgm.data_per_vertex[dpv_key]._lengths = tgm.streamlines._lengths for dps_key in init_as.data_per_streamline.keys(): dtype = init_as.data_per_streamline[dps_key].dtype @@ -1369,13 +1369,13 @@ def _initialize_empty_trx( # noqa: C901 logging.debug( f"Initializing {dps_key} (dps) with and dtype: {dtype.name}" ) - trx.data_per_streamline[dps_key] = _create_memmap( + tgm.data_per_streamline[dps_key] = _create_memmap( dps_filename, mode="w+", shape=shape, dtype=dtype ) - trx._uncompressed_folder_handle = tmp_dir + tgm._uncompressed_folder_handle = tmp_dir - return trx + return tgm def _create_trx_from_pointer( # noqa: C901 header: dict, @@ -1402,12 +1402,12 @@ def _create_trx_from_pointer( # noqa: C901 TrxFile A TrxFile constructed from the pointer provided. """ - trx = TrxFile() - trx.header = header + tgm = TrxFile() + tgm.header = header # Handle empty TRX files early - no positions/offsets to load if header["NB_STREAMLINES"] == 0 or header["NB_VERTICES"] == 0: - return trx + return tgm positions, offsets = None, None for elem_filename in dict_pointer_size.keys(): @@ -1440,23 +1440,23 @@ def _create_trx_from_pointer( # noqa: C901 # Parse/walk the directory tree if base == "positions" and folder == "": - if size != trx.header["NB_VERTICES"] * 3 or dim != 3: + if size != tgm.header["NB_VERTICES"] * 3 or dim != 3: raise ValueError("Wrong data size/dimensionality.") positions = _create_memmap( filename, mode="r+", offset=mem_adress, - shape=(trx.header["NB_VERTICES"], 3), + shape=(tgm.header["NB_VERTICES"], 3), dtype=ext[1:], ) elif base == "offsets" and folder == "": - if size != trx.header["NB_STREAMLINES"] + 1 or dim != 1: + if size != tgm.header["NB_STREAMLINES"] + 1 or dim != 1: raise ValueError("Wrong offsets size/dimensionality.") offsets = _create_memmap( filename, mode="r+", offset=mem_adress, - shape=(trx.header["NB_STREAMLINES"] + 1,), + shape=(tgm.header["NB_STREAMLINES"] + 1,), dtype=ext[1:], ) if offsets[-1] != 0: @@ -1464,23 +1464,23 @@ def _create_trx_from_pointer( # noqa: C901 else: lengths = [0] elif folder == "dps": - nb_scalar = size / trx.header["NB_STREAMLINES"] + nb_scalar = size / tgm.header["NB_STREAMLINES"] if not nb_scalar.is_integer() or nb_scalar != dim: raise ValueError("Wrong dps size/dimensionality.") else: - shape = (trx.header["NB_STREAMLINES"], int(nb_scalar)) + shape = (tgm.header["NB_STREAMLINES"], int(nb_scalar)) - trx.data_per_streamline[base] = _create_memmap( + tgm.data_per_streamline[base] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) elif folder == "dpv": - nb_scalar = size / trx.header["NB_VERTICES"] + nb_scalar = size / tgm.header["NB_VERTICES"] if not nb_scalar.is_integer() or nb_scalar != dim: raise ValueError("Wrong dpv size/dimensionality.") else: - shape = (trx.header["NB_VERTICES"], int(nb_scalar)) + shape = (tgm.header["NB_VERTICES"], int(nb_scalar)) - trx.data_per_vertex[base] = _create_memmap( + tgm.data_per_vertex[base] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) elif folder.startswith("dpg"): @@ -1492,9 +1492,9 @@ def _create_trx_from_pointer( # noqa: C901 # Handle the two-layers architecture data_name = os.path.basename(base) sub_folder = os.path.basename(folder) - if sub_folder not in trx.data_per_group: - trx.data_per_group[sub_folder] = {} - trx.data_per_group[sub_folder][data_name] = _create_memmap( + 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": @@ -1504,7 +1504,7 @@ def _create_trx_from_pointer( # noqa: C901 raise ValueError("Wrong group dimensionality.") else: shape = (int(size),) - trx.groups[base] = _create_memmap( + tgm.groups[base] = _create_memmap( filename, mode="r+", offset=mem_adress, shape=shape, dtype=ext[1:] ) else: @@ -1512,19 +1512,19 @@ def _create_trx_from_pointer( # noqa: C901 # All essential array must be declared if positions is not None and offsets is not None: - trx.streamlines._data = positions - trx.streamlines._offsets = offsets[:-1] - trx.streamlines._lengths = lengths + tgm.streamlines._data = positions + tgm.streamlines._offsets = offsets[:-1] + tgm.streamlines._lengths = lengths else: raise ValueError("Missing essential data.") - for dpv_key in trx.data_per_vertex: - tmp = trx.data_per_vertex[dpv_key] - trx.data_per_vertex[dpv_key] = ArraySequence() - trx.data_per_vertex[dpv_key]._data = tmp - trx.data_per_vertex[dpv_key]._offsets = offsets[:-1] - trx.data_per_vertex[dpv_key]._lengths = lengths - return trx + for dpv_key in tgm.data_per_vertex: + tmp = tgm.data_per_vertex[dpv_key] + tgm.data_per_vertex[dpv_key] = ArraySequence() + tgm.data_per_vertex[dpv_key]._data = tmp + tgm.data_per_vertex[dpv_key]._offsets = offsets[:-1] + tgm.data_per_vertex[dpv_key]._lengths = lengths + return tgm def resize( # noqa: C901 self, @@ -1572,7 +1572,7 @@ def resize( # noqa: C901 logging.debug("TrxFile of the right size, no resizing.") return - trx = self._initialize_empty_trx(nb_streamlines, nb_vertices, init_as=self) + tgm = self._initialize_empty_trx(nb_streamlines, nb_vertices, init_as=self) logging.info( "Resizing streamlines from size " @@ -1586,11 +1586,11 @@ def resize( # noqa: C901 # Copy the fixed-sized info from the original TrxFile to the new # (resized) one. if nb_streamlines < self.header["NB_STREAMLINES"]: - trx._copy_fixed_arrays_from(self, nb_strs_to_copy=nb_streamlines) + tgm._copy_fixed_arrays_from(self, nb_strs_to_copy=nb_streamlines) else: - trx._copy_fixed_arrays_from(self) + tgm._copy_fixed_arrays_from(self) - tmp_dir = trx._uncompressed_folder_handle.name + tmp_dir = tgm._uncompressed_folder_handle.name if len(self.groups.keys()) > 0: os.mkdir(os.path.join(tmp_dir, "groups/")) @@ -1603,15 +1603,15 @@ def resize( # noqa: C901 # Remove groups indices if resizing down tmp = self.groups[group_key][self.groups[group_key] < strs_end] - trx.groups[group_key] = _create_memmap( + tgm.groups[group_key] = _create_memmap( group_name, mode="w+", shape=(len(tmp),), dtype=group_dtype ) logging.debug(f"{group_key} group went from {ori_len} items to {len(tmp)}") - trx.groups[group_key][:] = tmp + tgm.groups[group_key][:] = tmp if delete_dpg: self.close() - self.__dict__ = trx.__dict__ + self.__dict__ = tgm.__dict__ return if len(self.data_per_group.keys()) > 0: @@ -1619,8 +1619,8 @@ def resize( # noqa: C901 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 group_key not in trx.data_per_group: - trx.data_per_group[group_key] = {} + if group_key not in tgm.data_per_group: + tgm.data_per_group[group_key] = {} for dpg_key in self.data_per_group[group_key].keys(): dpg_dtype = self.data_per_group[group_key][dpg_key].dtype @@ -1630,18 +1630,18 @@ def resize( # noqa: C901 ) shape = self.data_per_group[group_key][dpg_key].shape - if dpg_key not in trx.data_per_group[group_key]: - trx.data_per_group[group_key][dpg_key] = {} - trx.data_per_group[group_key][dpg_key] = _create_memmap( + if dpg_key not in tgm.data_per_group[group_key]: + tgm.data_per_group[group_key][dpg_key] = {} + tgm.data_per_group[group_key][dpg_key] = _create_memmap( dpg_filename, mode="w+", shape=shape, dtype=dpg_dtype ) - trx.data_per_group[group_key][dpg_key][:] = self.data_per_group[ + tgm.data_per_group[group_key][dpg_key][:] = self.data_per_group[ group_key ][dpg_key] self.close() - self.__dict__ = trx.__dict__ + self.__dict__ = tgm.__dict__ def get_dtype_dict(self): """Get the dtype dictionary for the TrxFile. @@ -1711,20 +1711,20 @@ def append(self, obj, extra_buffer: int = 0) -> None: self._append_trx(obj, extra_buffer=extra_buffer) - def _append_trx(self, trx: Type["TrxFile"], extra_buffer: int = 0) -> None: + def _append_trx(self, tgm: Type["TrxFile"], extra_buffer: int = 0) -> None: """Append a TrxFile to another (with buffer support). Parameters ---------- - trx : TrxFile + tgm : TrxFile The TrxFile to append to the current TrxFile. extra_buffer : int, optional The additional buffer space required to append data. Default is 0. """ strs_end, pts_end = self._get_real_len() - nb_streamlines = strs_end + trx.header["NB_STREAMLINES"] - nb_vertices = pts_end + trx.header["NB_VERTICES"] + nb_streamlines = strs_end + tgm.header["NB_STREAMLINES"] + nb_vertices = pts_end + tgm.header["NB_VERTICES"] if ( self.header["NB_STREAMLINES"] < nb_streamlines @@ -1734,7 +1734,7 @@ def _append_trx(self, trx: Type["TrxFile"], extra_buffer: int = 0) -> None: nb_streamlines=nb_streamlines + extra_buffer, nb_vertices=nb_vertices + extra_buffer * 100, ) - _ = concatenate([self, trx], preallocation=True, delete_groups=True) + _ = concatenate([self, tgm], preallocation=True, delete_groups=True) def get_group( self, key: str, keep_group: bool = True, copy_safe: bool = False @@ -1972,10 +1972,10 @@ def from_sft(sft, dtype_dict=None): "dtype." ) - trx = TrxFile( + tgm = TrxFile( nb_vertices=len(sft.streamlines._data), nb_streamlines=len(sft.streamlines) ) - trx.header = { + tgm.header = { "DIMENSIONS": sft.dimensions.tolist(), "VOXEL_TO_RASMM": sft.affine.tolist(), "NB_VERTICES": len(sft.streamlines._data), @@ -1995,13 +1995,13 @@ def from_sft(sft, dtype_dict=None): tmp_streamlines._offsets = tmp_streamlines._offsets.astype(offsets_dtype) tmp_streamlines._data = tmp_streamlines._data.astype(positions_dtype) - trx.streamlines = tmp_streamlines + tgm.streamlines = tmp_streamlines for key in sft.data_per_point: dtype_to_use = ( dtype_dict["dpv"][key] if key in dtype_dict["dpv"] else np.float32 ) - trx.data_per_vertex[key] = sft.data_per_point[key] - trx.data_per_vertex[key]._data = sft.data_per_point[key]._data.astype( + tgm.data_per_vertex[key] = sft.data_per_point[key] + tgm.data_per_vertex[key]._data = sft.data_per_point[key]._data.astype( dtype_to_use ) @@ -2009,22 +2009,22 @@ def from_sft(sft, dtype_dict=None): dtype_to_use = ( dtype_dict["dps"][key] if key in dtype_dict["dps"] else np.float32 ) - trx.data_per_streamline[key] = sft.data_per_streamline[key].astype( + tgm.data_per_streamline[key] = sft.data_per_streamline[key].astype( dtype_to_use ) # For safety and for RAM, convert the whole object to memmaps tmp_dir = get_trx_tmp_dir() - save(trx, tmp_dir.name) - trx.close() - trx = load_from_directory(tmp_dir.name) - trx._uncompressed_folder_handle = tmp_dir + save(tgm, tmp_dir.name) + tgm.close() + tgm = load_from_directory(tmp_dir.name) + tgm._uncompressed_folder_handle = tmp_dir sft.to_space(old_space) sft.to_origin(old_origin) del tmp_streamlines - return trx + return tgm @staticmethod def from_tractogram( @@ -2073,13 +2073,13 @@ def from_tractogram( "dtype." ) - trx = TrxFile( + tgm = TrxFile( nb_vertices=len(tractogram.streamlines._data), nb_streamlines=len(tractogram.streamlines), ) affine, dimensions, _, _ = get_reference_info_wrapper(reference) - trx.header = { + tgm.header = { "DIMENSIONS": dimensions, "VOXEL_TO_RASMM": affine, "NB_VERTICES": len(tractogram.streamlines._data), @@ -2092,13 +2092,13 @@ def from_tractogram( tmp_streamlines._offsets = tmp_streamlines._offsets.astype(offsets_dtype) tmp_streamlines._data = tmp_streamlines._data.astype(positions_dtype) - trx.streamlines = tmp_streamlines + tgm.streamlines = tmp_streamlines for key in tractogram.data_per_point: dtype_to_use = ( dtype_dict["dpv"][key] if key in dtype_dict["dpv"] else np.float32 ) - trx.data_per_vertex[key] = tractogram.data_per_point[key] - trx.data_per_vertex[key]._data = tractogram.data_per_point[ + tgm.data_per_vertex[key] = tractogram.data_per_point[key] + tgm.data_per_vertex[key]._data = tractogram.data_per_point[ key ]._data.astype(dtype_to_use) @@ -2106,19 +2106,19 @@ def from_tractogram( dtype_to_use = ( dtype_dict["dps"][key] if key in dtype_dict["dps"] else np.float32 ) - trx.data_per_streamline[key] = tractogram.data_per_streamline[key].astype( + tgm.data_per_streamline[key] = tractogram.data_per_streamline[key].astype( dtype_to_use ) # For safety and for RAM, convert the whole object to memmaps tmp_dir = get_trx_tmp_dir() - save(trx, tmp_dir.name) - trx.close() + save(tgm, tmp_dir.name) + tgm.close() - trx = load_from_directory(tmp_dir.name) + tgm = load_from_directory(tmp_dir.name) del tmp_streamlines - return trx + return tgm def to_tractogram(self, resize=False): """Convert this TrxFile to a nibabel Tractogram. diff --git a/trx/workflows.py b/trx/workflows.py index f726c2e7..d5a23b77 100644 --- a/trx/workflows.py +++ b/trx/workflows.py @@ -124,8 +124,8 @@ def convert_dsi_studio( if out_ext != ".trx": save_tractogram(sft_flip, out_tractogram, bbox_valid_check=not keep_invalid) else: - trx = tmm.TrxFile.from_sft(sft_flip) - tmm.save(trx, out_tractogram) + tgm = tmm.TrxFile.from_sft(sft_flip) + tmm.save(tgm, out_tractogram) def convert_tractogram( # noqa: C901 @@ -169,9 +169,9 @@ def convert_tractogram( # noqa: C901 if in_ext != ".trx": sft = load_sft_with_reference(in_tractogram, reference, bbox_check=False) else: - trx = tmm.load(in_tractogram) - sft = trx.to_sft() - trx.close() + tgm = tmm.load(in_tractogram) + sft = tgm.to_sft() + tgm.close() if out_ext != ".trx": if out_ext == ".vtk": @@ -185,13 +185,13 @@ def convert_tractogram( # noqa: C901 ) save_tractogram(sft, out_tractogram, bbox_valid_check=False) else: - trx = tmm.TrxFile.from_sft(sft) - if trx.streamlines._data.dtype.name != pos_dtype: - trx.streamlines._data = trx.streamlines._data.astype(pos_dtype) - if trx.streamlines._offsets.dtype.name != offsets_dtype: - trx.streamlines._offsets = trx.streamlines._offsets.astype(offsets_dtype) - tmm.save(trx, out_tractogram) - trx.close() + tgm = tmm.TrxFile.from_sft(sft) + if tgm.streamlines._data.dtype.name != pos_dtype: + tgm.streamlines._data = tgm.streamlines._data.astype(pos_dtype) + if tgm.streamlines._offsets.dtype.name != offsets_dtype: + tgm.streamlines._offsets = tgm.streamlines._offsets.astype(offsets_dtype) + tmm.save(tgm, out_tractogram) + tgm.close() def tractogram_simple_compare(in_tractograms, reference): @@ -697,9 +697,9 @@ def generate_trx_from_scratch( # noqa: C901 for arg in dpg: _write_data_array(tmp_dir_name, "dpg", arg, is_dpg=True) - trx = tmm.load(tmp_dir_name) - tmm.save(trx, out_tractogram) - trx.close() + tgm = tmm.load(tmp_dir_name) + tmm.save(tgm, out_tractogram) + tgm.close() def manipulate_trx_datatype(in_filename, out_filename, dict_dtype): # noqa: C901 @@ -719,7 +719,7 @@ def manipulate_trx_datatype(in_filename, out_filename, dict_dtype): # noqa: C90 None Writes the converted TRX to ``out_filename``. """ - trx = tmm.load(in_filename) + tgm = tmm.load(in_filename) # For each key in dict_dtype, we create a new memmap with the new dtype # and we copy the data from the old memmap to the new one. @@ -727,53 +727,53 @@ def manipulate_trx_datatype(in_filename, out_filename, dict_dtype): # noqa: C90 for key in dict_dtype: if key == "positions": tmp_mm = _create_temp_memmap( - tmp_dir_name, dict_dtype[key], trx.streamlines._data.shape + tmp_dir_name, dict_dtype[key], tgm.streamlines._data.shape ) - tmp_mm[:] = trx.streamlines._data[:] - trx.streamlines._data = tmp_mm + tmp_mm[:] = tgm.streamlines._data[:] + tgm.streamlines._data = tmp_mm elif key == "offsets": tmp_mm = _create_temp_memmap( - tmp_dir_name, dict_dtype[key], trx.streamlines._offsets.shape + tmp_dir_name, dict_dtype[key], tgm.streamlines._offsets.shape ) - tmp_mm[:] = trx.streamlines._offsets[:] - trx.streamlines._offsets = tmp_mm + tmp_mm[:] = tgm.streamlines._offsets[:] + tgm.streamlines._offsets = tmp_mm elif key == "dpv": for key_dpv in dict_dtype[key]: tmp_mm = _create_temp_memmap( tmp_dir_name, dict_dtype[key][key_dpv], - trx.data_per_vertex[key_dpv]._data.shape, + tgm.data_per_vertex[key_dpv]._data.shape, ) - tmp_mm[:] = trx.data_per_vertex[key_dpv]._data[:] - trx.data_per_vertex[key_dpv]._data = tmp_mm + tmp_mm[:] = tgm.data_per_vertex[key_dpv]._data[:] + tgm.data_per_vertex[key_dpv]._data = tmp_mm elif key == "dps": for key_dps in dict_dtype[key]: tmp_mm = _create_temp_memmap( tmp_dir_name, dict_dtype[key][key_dps], - trx.data_per_streamline[key_dps].shape, + tgm.data_per_streamline[key_dps].shape, ) - tmp_mm[:] = trx.data_per_streamline[key_dps][:] - trx.data_per_streamline[key_dps] = tmp_mm + tmp_mm[:] = tgm.data_per_streamline[key_dps][:] + tgm.data_per_streamline[key_dps] = tmp_mm elif key == "dpg": for key_group in dict_dtype[key]: for key_dpg in dict_dtype[key][key_group]: tmp_mm = _create_temp_memmap( tmp_dir_name, dict_dtype[key][key_group][key_dpg], - trx.data_per_group[key_group][key_dpg].shape, + tgm.data_per_group[key_group][key_dpg].shape, ) - tmp_mm[:] = trx.data_per_group[key_group][key_dpg][:] - trx.data_per_group[key_group][key_dpg] = tmp_mm + tmp_mm[:] = tgm.data_per_group[key_group][key_dpg][:] + tgm.data_per_group[key_group][key_dpg] = tmp_mm elif key == "groups": for key_group in dict_dtype[key]: tmp_mm = _create_temp_memmap( tmp_dir_name, dict_dtype[key][key_group], - trx.groups[key_group].shape, + tgm.groups[key_group].shape, ) - tmp_mm[:] = trx.groups[key_group][:] - trx.groups[key_group] = tmp_mm + tmp_mm[:] = tgm.groups[key_group][:] + tgm.groups[key_group] = tmp_mm - tmm.save(trx, out_filename) - trx.close() + tmm.save(tgm, out_filename) + tgm.close()