NathanHowell/jax-digest

T-digest quantile estimation implemented in JAX

★ 0Forks 0PythonGitHub ↗Compare
jax

README

jax-digest

T-digest quantile estimation implemented in JAX. All operations are jax.jit compatible.

T-digest is a data structure for accurate online estimation of quantiles, particularly at the tails of distributions. This implementation uses fixed-capacity centroid buffers and jax.lax.scan for compression, making it fully compatible with JAX's JIT compilation and functional transformation model.

Install

uv add jax-digest

Usage

import jax_digest

# Create a digest
td = jax_digest.new(delta=100)

# Add values one at a time
for x in data:
    td = jax_digest.update(td, x)

# Query quantiles
p50 = jax_digest.quantile(td, 0.5)
p99 = jax_digest.quantile(td, 0.99)

# Query CDF
fraction_below = jax_digest.cdf(td, 42.0)

# Trimmed mean
tm = jax_digest.trimmed_mean(td, 0.1, 0.9)

# Merge two digests
merged = jax_digest.merge(td1, td2)

API

Function Description
new(delta=100, capacity=None) Create an empty t-digest. delta controls accuracy vs memory (higher = more centroids).
update(td, value) Add a single value. Auto-compresses when the buffer fills.
batch_update(td, values) Add an array of values.
merge(td1, td2) Merge two t-digests.
quantile(td, q) Estimate the value at quantile q (0 to 1).
quantiles(td, qs) Vectorized quantile estimation for an array of quantile values.
cdf(td, value) Estimate the CDF (fraction of values <= value).
trimmed_mean(td, lower_q, upper_q) Estimate the mean between two quantiles.

All query and update functions are decorated with @jax.jit.

Implementation details

  • Fixed-capacity centroid buffer with static shapes for JIT compatibility
  • k1 (arcsin) scale function for high accuracy in distribution tails
  • Alternating merge direction to avoid systematic bias
  • Compression via jax.lax.scan — no Python-level loops in the hot path

License

BSD Zero Clause License. See LICENSE.

Contributors

NathanHowell

Issues