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.
uv add jax-digest
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)| 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.
- 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
BSD Zero Clause License. See LICENSE.