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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion docs/source/dev.md
Original file line number Diff line number Diff line change
Expand Up @@ -251,7 +251,7 @@ def load(filename, reference=None):
Examples
--------
>>> from trx.io import load
>>> trx = load("tractogram.trx")
>>> tgm = load("tractogram.trx")
"""
pass
```
Expand Down
50 changes: 25 additions & 25 deletions examples/plot_dps_dpv.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,22 +52,22 @@
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)
# -------------------------------
#
# 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}")
Expand All @@ -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]
Expand All @@ -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}")
Expand All @@ -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}':")
Expand Down Expand Up @@ -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")

Expand All @@ -177,23 +177,23 @@

# 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
# -------
#
# 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
Expand Down
48 changes: 24 additions & 24 deletions examples/plot_groups.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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")

# %%
Expand All @@ -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")
Expand All @@ -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}'")

# %%
Expand All @@ -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}")

# %%
Expand Down Expand Up @@ -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}")

Expand All @@ -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
Expand Down
38 changes: 19 additions & 19 deletions examples/plot_read_write_trx.py
Original file line number Diff line number Diff line change
Expand Up @@ -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!")

Expand All @@ -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
Expand All @@ -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]}")

Expand All @@ -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
Expand All @@ -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")

# %%
Expand All @@ -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")

Expand All @@ -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
Expand Down
10 changes: 5 additions & 5 deletions trx/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -233,15 +233,15 @@ 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,
delete_groups=delete_groups or not has_group,
check_space_attributes=True,
preallocation=False,
)
save(trx, str(out_tractogram))
save(tgm, str(out_tractogram))

typer.echo(
typer.style(
Expand Down Expand Up @@ -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:")
Expand All @@ -1021,7 +1021,7 @@ def info(
typer.echo("--------- -------")
typer.echo(f"{total_uncompressed:>9} {num_files} files")

trx.close()
tgm.close()


def main():
Expand Down
Loading
Loading