Skip to main content

phasecurvefit: Construct Paths through Phase-Space Points

PyPI version Python versions

Construct paths through phase-Space points, supporting many different algorithms.

Features

  • JAX-powered: Fully compatible with JAX transformations (jit, vmap, grad)
  • GPU-ready: Runs on CPU, GPU, or TPU via JAX
  • Type-safe: Comprehensive (optionally runtime checked) type hints with jaxtyping
  • Pluggable metrics: Customizable distance metrics for different physical interpretations
  • Pluggable query strategies: Flexible neighbor search strategies (e.g., brute-force, KD-tree) to optimize performance
  • Pluggable orderers: One interface over multiple ordering algorithms — the velocity-following walk and an MST backbone for near-closed loops
  • Highly customizable ML setup and training: Well-chosen defaults with highly flexible customization for specific use-cases.
  • Physical units: Optional support via unxt for unit-aware calculations

Installation

Install the core package:

pip install phasecurvefit[all]

Or with uv:

uv add phasecurvefit[all]
from source, using uv
uv add git+https://github.com/GalacticDynamics/phasecurvefit.git@main

You can customize the branch by replacing main with any other branch name.

building from source
cd /path/to/parent
git clone https://github.com/GalacticDynamics/phasecurvefit.git
cd phasecurvefit
uv pip install -e .  # editable mode

Optional Dependencies

phasecurvefit has optional dependencies for extended functionality:

  • unxt: Physical units support for phase-space calculations
  • tree (jaxkd): Spatial KD-tree queries for large datasets

Install with optional dependencies:

# pip install phasecurvefit[all]  # Install with all extras
pip install phasecurvefit[interop]  # Install with unxt for unit support
pip install phasecurvefit[kdtree]  # Install with jaxkd for KD-tree strategy

Or with uv:

# uv add phasecurvefit --extra all  # installs all extras
uv add phasecurvefit --extra interop
uv add phasecurvefit --extra kdtree

Quick Start

import jax
import jax.numpy as jnp
import phasecurvefit as pcf

# Create phase-space observations as dictionaries (Cartesian coordinates)
pos = {
    "x": jnp.array([0.0, 1.0, 2.0, 3.0, 4.0]),
    "y": jnp.array([0.0, 0.5, 1.0, 1.5, 2.0]),
}
vel = {
    "x": jnp.array([1.0, 1.0, 1.0, 1.0, 1.0]),
    "y": jnp.array([0.5, 0.5, 0.5, 0.5, 0.5]),
}

# Step 1: Order the observations (use KD-tree for spatial neighbor prefiltering)
config = pcf.WalkConfig(strategy=pcf.strats.KDTree(k=3))  # k=3 for this small dataset
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config))
print(result.indices)  # Initial ordering

# Step 2: Create normalizer and autoencoder
key = jax.random.key(0)
normalizer = pcf.nn.StandardScalerNormalizer(pos, vel)
autoencoder = pcf.nn.PathAutoencoder.make(
    normalizer, gamma_range=result.gamma_range, key=key
)

# Step 3: Configure and run training
train_config = pcf.nn.TrainingConfig(
    n_epochs_encoder=100,  # Encoder-only epochs
    n_epochs_both=50,  # Joint training epochs
    show_pbar=False,  # Disable progress bar
)

# Train the autoencoder
result, _, losses = pcf.nn.train_autoencoder(
    autoencoder, result, config=train_config, key=key
)

print(result.indices)  # Post-training ordering

With Physical Units

When unxt is installed, you can use physical units throughout the workflow:

import jax
import jax.numpy as jnp
import phasecurvefit as pcf
import unxt as u

# Create phase-space observations with units
pos = {
    "x": u.Q([0.0, 1.0, 2.0, 3.0, 4.0], "kpc"),
    "y": u.Q([0.0, 0.5, 1.0, 1.5, 2.0], "kpc"),
}
vel = {
    "x": u.Q([1.0, 1.0, 1.0, 1.0, 1.0], "km/s"),
    "y": u.Q([0.5, 0.5, 0.5, 0.5, 0.5], "km/s"),
}

# Step 1: Order with units (units are preserved throughout)
metric_scale = u.Q(1.0, "kpc")
config = pcf.WalkConfig(strategy=pcf.strats.KDTree(k=3))
result = pcf.order(
    pos,
    vel,
    pcf.orderers.LocalFlowOrderer(config=config, metric_scale=metric_scale),
    metadata=pcf.StateMetadata(usys=u.unitsystems.galactic),
)

# Step 2: Create normalizer and autoencoder (handles units automatically)
key = jax.random.key(0)
normalizer = pcf.nn.StandardScalerNormalizer(pos, vel)
autoencoder = pcf.nn.PathAutoencoder.make(
    normalizer, gamma_range=result.gamma_range, key=key
)
result, _, losses = pcf.nn.train_autoencoder(
    autoencoder, result, config=train_config, key=key
)

Orderers

The ordering step is pluggable. Every orderer implements the same interface — order(positions, velocities) — and returns an OrderingResult that feeds the autoencoder unchanged, so orderers are interchangeable:

  • LocalFlowOrderer — the velocity-following walk (wraps walk_local_flow). Follows a coherent flow from a start point.
  • MSTOrderer — a minimum-spanning-tree backbone. It needs no start point (the graph diameter finds the two tips itself), which makes it ideal for near-closed loops where the velocity field reverses and a single walk covers only one arm. Requires the mst extra.
import jax.numpy as jnp
import phasecurvefit as pcf

# Points along a curve
t = jnp.linspace(0.0, 1.0, 60)
pos = {"x": 10.0 * t, "y": jnp.sin(3.0 * t)}
vel = {"x": jnp.ones(60), "y": 3.0 * jnp.cos(3.0 * t)}

# The velocity-following walk, via the orderer interface
walk_orderer = pcf.orderers.LocalFlowOrderer(metric_scale=1.0, start_idx=0)
walk_result = walk_orderer.order(pos, vel)

# The MST backbone (no start point needed)
mst_orderer = pcf.orderers.MSTOrderer(k=8, jump_cap=2.0)
mst_result = pcf.order(pos, vel, mst_orderer)  # or mst_orderer.order(pos, vel)

# Either result feeds the autoencoder unchanged
print(mst_result.gamma_range)  # (-1.0, 1.0)

MSTOrderer also has opt-in velocity mechanisms (velocity_weight, sever_cos_threshold, orient_by_velocity) for self-overlapping streams. See the Orderers Guide and the Migration Guide.

Distance Metrics

The algorithm supports pluggable distance metrics to control how points are ordered. The default metric is AlignedMomentumDistanceMetric, which combines spatial proximity with velocity alignment:

import jax.numpy as jnp
import phasecurvefit as pcf

# Define simple Cartesian arrays (not quantities)
pos = {"x": jnp.array([0.0, 1.0, 2.0]), "y": jnp.array([0.0, 0.5, 1.0])}
vel = {"x": jnp.array([1.0, 1.0, 1.0]), "y": jnp.array([0.5, 0.5, 0.5])}

# Use default metric (AlignedMomentumDistanceMetric)
config = pcf.WalkConfig()
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config))

Using Different Metrics

phasecurvefit provides three built-in metrics:

  1. AlignedMomentumDistanceMetric (default): Combines spatial distance with velocity alignment (momentum-weighted nearest neighbor)
  2. FullPhaseSpaceDistanceMetric: True 6D Euclidean distance in phase space
  3. SpatialDistanceMetric: Pure spatial distance, ignoring velocity
import jax.numpy as jnp
import phasecurvefit as pcf

# Define simple Cartesian arrays (not quantities)
pos = {"x": jnp.array([0.0, 1.0, 2.0]), "y": jnp.array([0.0, 0.5, 1.0])}
vel = {"x": jnp.array([1.0, 1.0, 1.0]), "y": jnp.array([0.5, 0.5, 0.5])}

# Pure spatial ordering (ignores velocity)
config_spatial = pcf.WalkConfig(metric=pcf.metrics.SpatialDistanceMetric())
result = pcf.order(
    pos, vel, pcf.orderers.LocalFlowOrderer(config=config_spatial, metric_scale=0.0)
)

# Full 6D phase-space distance
config_phase = pcf.WalkConfig(metric=pcf.metrics.FullPhaseSpaceDistanceMetric())
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config_phase))

Custom Metrics

You can define custom metrics by subclassing AbstractDistanceMetric:

import jax
import jax.numpy as jnp
import phasecurvefit as pcf


class WeightedPhaseSpaceMetric(pcf.metrics.AbstractDistanceMetric):
    """Custom weighted phase-space metric."""

    def __call__(self, current_pos, current_vel, positions, velocities, metric_scale):
        # Compute position distance
        pos_diff = jax.tree.map(jnp.subtract, positions, current_pos)
        pos_dist_sq = sum(jax.tree.leaves(jax.tree.map(jnp.square, pos_diff)))

        # Compute velocity distance
        vel_diff = jax.tree.map(jnp.subtract, velocities, current_vel)
        vel_dist_sq = sum(jax.tree.leaves(jax.tree.map(jnp.square, vel_diff)))

        # Custom weighting scheme
        return jnp.sqrt(pos_dist_sq + (metric_scale**2) * vel_dist_sq)


# Use custom metric via WalkConfig
config = pcf.WalkConfig(metric=WeightedPhaseSpaceMetric())
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config))

See the Metrics Guide for more details and examples.

Query Strategies

The algorithm supports pluggable query strategies to control how neighbors are found. A strategy determines which points are considered as potential next steps in the walk.

phasecurvefit provides two built-in strategies:

  1. BruteForce (default): Compute distances to all remaining points and select the nearest one. Efficient for small to medium datasets.
  2. KDTree: Use spatial KD-tree prefiltering to accelerate neighbor searches for large datasets (requires optional jaxkd dependency).

Using Built-in Strategies

import jax.numpy as jnp
import phasecurvefit as pcf

# Define simple Cartesian arrays
pos = {"x": jnp.array([0.0, 1.0, 2.0]), "y": jnp.array([0.0, 0.5, 1.0])}
vel = {"x": jnp.array([1.0, 1.0, 1.0]), "y": jnp.array([0.5, 0.5, 0.5])}

# Default strategy (brute-force — no configuration needed)
config_brute = pcf.WalkConfig(strategy=pcf.strats.BruteForce())
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config_brute))

# KD-tree strategy for faster neighbor queries (large datasets)
config_kdtree = pcf.WalkConfig(strategy=pcf.strats.KDTree(k=2))
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config_kdtree))

Custom Query Strategies

You can define custom strategies by subclassing AbstractQueryStrategy:

import jax.numpy as jnp
import phasecurvefit as pcf


class SmallestIndexStrategy(pcf.strats.AbstractQueryStrategy):
    """Custom strategy: select the smallest unvisited index.

    This is a toy example showing how to implement a custom strategy.
    By returning uniform distances, argmin selects the smallest index
    deterministically. In practice, distance-based strategies like BruteForce
    are more useful.
    """

    def init(self, positions, /, *, metadata):
        """No persistent state needed."""
        return None

    def query(
        self,
        state,
        /,
        current_pos,
        current_vel,
        positions,
        velocities,
        metric_fn,
        metric_scale,
    ):
        """Return uniform distances to all points.

        Since all distances are equal, the walk algorithm's argmin will
        deterministically select the smallest unvisited index.
        """
        # Get number of points
        n_points = len(next(iter(positions.values())))

        # Return uniform distances to all points
        # argmin will pick the smallest unvisited index
        distances = jnp.ones(n_points)

        return pcf.strats.QueryResult(distances=distances, indices=None)


# Use custom strategy via WalkConfig
config = pcf.WalkConfig(strategy=SmallestIndexStrategy())
result = pcf.order(pos, vel, pcf.orderers.LocalFlowOrderer(config=config))

AI Usage Disclosure

Portions of this codebase (including tests and documentation) were refactored and generated with the assistance of Language Models. All AI contributions have been and will continue to be reviewed and verified by the human maintainers.

Release files for phasecurvefit 0.3.2

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for phasecurvefit 0.3.2
File Size Uploaded
phasecurvefit-0.3.2.tar.gz 3.8 MB Details

Built distribution (wheel)

Table of built distributions (wheels) for phasecurvefit 0.3.2
File Interpreter ABI Platform
phasecurvefit-0.3.2-py3-none-any.whl Python 3 none any Details

Total release size: 3.9 MB

Release files / phasecurvefit-0.3.2.tar.gz

Download URL phasecurvefit-0.3.2.tar.gz
Size 3.8 MB
Tags Source
SHA-256 checksum
How to use checksums
10dfeead7eff255d97a97598333a6f1e440e1042cf1403a4e93f4c2e3d5921a5
BLAKE2b-256 checksum
How to use checksums
4c9dadcc17fa9c06a369762f9a126206747a4e92f04e18e9745a6c405ebac922
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 24, 2026.

Transparency log

Release files / phasecurvefit-0.3.2-py3-none-any.whl

Download URL phasecurvefit-0.3.2-py3-none-any.whl
Size 99.4 kB
Tags Python 3
SHA-256 checksum
How to use checksums
8eaa054c57f551a3d26541eb2e914ea47325d874a9ec3a3770e4e198ecd924c6
BLAKE2b-256 checksum
How to use checksums
a7197547d51d976b189b88237ba6b865b045d01cca40f2e3d390ebd9fcacea16
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 24, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.3.2 This release

2 release files

0.3.1

2 release files

0.3.0

2 release files

0.2.0

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page