Construct paths through phase-Space points, supporting many different algorithms.
Many datasets are samples along a curve in phase space whose order along the curve is unknown. Before fitting a model to such a curve you need two things: an ordering coordinate for every sample, and a smooth track through them. Doing this by hand, or with a position-only nearest-neighbor or clustering method, breaks down in exactly the cases that matter:
- Curves that cross or fold back on themselves. Where two strands meet, the nearest point is often on the wrong strand. phasecurvefit uses velocities as well as positions, so the ordering stays on the right strand (see the epitrochoid tutorials).
- No known starting point. The MST orderer finds the two ends of the curve itself, so no progenitor position or hand-picked start index is needed (MST tutorial).
- Incomplete orderings. A conservative walk orders a reliable subset; an autoencoder then assigns an ordering coordinate γ to every sample and learns a smooth mean track through them (stream autoencoder tutorial).
- Contamination. A stream-plus-background mixture model gives each sample a calibrated membership probability, so interlopers can be down-weighted or removed (outlier-rejection tutorial).
- Use inside larger models. phasecurvefit is built on JAX: the walk, the
distance metrics and the neural networks work with
jit,vmapandgradand run on CPU or GPU. A training-free running-mean track is available when speed matters more than accuracy, for example inside a likelihood evaluated at every step of an MCMC (running-mean tutorial).
phasecurvefit is a reusable, tested library for momentum-weighted ordering, with
alternative orderers, gap filling, outlier rejection and optional physical units
(via unxt). It was built for stellar streams but applies to any ordered
phase-space data.
- 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
unxtfor unit-aware calculations
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@mainYou 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 modephasecurvefit 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 strategyOr with uv:
# uv add phasecurvefit --extra all # installs all extras
uv add phasecurvefit --extra interop
uv add phasecurvefit --extra kdtreeThe
tutorial notebooks
need packages beyond the runtime [all] extra — matplotlib for plotting and
galax for the mock-stream examples. Install them with the tutorials extra:
pip install phasecurvefit[tutorials]uv add phasecurvefit --extra tutorialsNote [all] intentionally does not include tutorials: all covers optional
runtime functionality, while tutorials covers packages only needed to run
the example notebooks.
phasecurvefit runs on GPU through JAX, but a plain pip install jax (what
phasecurvefit depends on) only ships a CPU-only jaxlib. If you have an NVIDIA
GPU and see:
An NVIDIA GPU may be present on this machine, but a CUDA-enabled jaxlib is
not installed. Falling back to cpu.
Install JAX's CUDA-enabled build alongside phasecurvefit:
pip install --upgrade "phasecurvefit[all]" "jax[cuda12]"Or with uv:
uv add phasecurvefit --extra all
uv add "jax[cuda12]"This pulls in self-contained NVIDIA CUDA/cuDNN wheels — you don't need the CUDA
toolkit installed system-wide — but you do still need a
compatible NVIDIA driver
for your GPU. --upgrade ensures pip actually swaps in the CUDA-enabled
jaxlib even if a CPU-only one is already installed. See the
JAX GPU installation guide
for other CUDA versions or platforms (TPU, ROCm), and verify the install with:
import jax
print(jax.devices()) # should list a CudaDevice, not just CpuDeviceimport 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 orderingWhen 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
)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 (wrapswalk_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 themstextra.
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.
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))phasecurvefit provides three built-in metrics:
- AlignedMomentumDistanceMetric (default): Combines spatial distance with velocity alignment (momentum-weighted nearest neighbor)
- FullPhaseSpaceDistanceMetric: True 6D Euclidean distance in phase space
- 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))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.
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:
- BruteForce (default): Compute distances to all remaining points and select the nearest one. Efficient for small to medium datasets.
- KDTree: Use spatial KD-tree prefiltering to accelerate neighbor searches
for large datasets (requires optional
jaxkddependency).
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))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))If you use phasecurvefit in published work, please cite the package via its
DOI, together with the paper behind whichever component you used.
component papers
- momentum-weighted ordering — Nibauer et al. (2022), arXiv:2201.12042
- SOM ordering — Starkman et al. (2023), MNRAS 522, 5022, arXiv:2212.00949
- mixture-model membership / outlier rejection — Hogg, Bovy & Lang (2010), arXiv:1008.4686
Machine-readable metadata for all of these is in
CITATION.cff;
BibTeX entries are in the documentation.
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.