From a96c5aa26a769a84a4d4b219e0c969ea6c8cf123 Mon Sep 17 00:00:00 2001 From: ryux1 Date: Sun, 6 Sep 2026 03:56:21 +0200 Subject: [PATCH] feat(tdigest): Add batch quantile queries Co-authored-by: tison --- CHANGELOG.md | 2 + benchmarks/tdigest/query.rs | 21 ++ datasketches/src/tdigest/sketch.rs | 181 ++++++++++++++---- .../tests/tdigest_test/sketch.rs | 32 ++++ 4 files changed, 198 insertions(+), 38 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 29ceccf7..034a0bd9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,7 @@ All significant changes to this project will be documented in this file. ### New features * Add KLL sketches behind the `kll` feature, with rank, quantile, PMF, and CDF queries, merging, totally ordered custom item types, a `KllFloat` adapter for non-NaN floating-point values, and serialization. +* Add `TDigestMut::quantiles` and `TDigest::quantiles` for querying several ranks in one centroid scan. ### Improvements @@ -18,6 +19,7 @@ All significant changes to this project will be documented in this file. * Improve truncated-input diagnostics across sketch deserializers. * Improve hash-backed sketch update performance for integer and raw-byte inputs. * Improve Bloom filter membership-and-insert performance and simplify Theta-family hash table thresholds. +* T-Digest batch quantile queries reuse one traversal for ranks supplied in nondecreasing order. ### Bug fixes diff --git a/benchmarks/tdigest/query.rs b/benchmarks/tdigest/query.rs index 0981657d..62422d3f 100644 --- a/benchmarks/tdigest/query.rs +++ b/benchmarks/tdigest/query.rs @@ -56,3 +56,24 @@ fn quantiles_2_sequential(bencher: Bencher) { ] }); } + +#[divan::bench] +fn quantiles_6_sequential(bencher: Bencher) { + let digest = prepared_digest(); + let ranks = [0.25, 0.5, 0.75, 0.9, 0.95, 0.99]; + + bencher + .bench_local(|| black_box(ranks).map(|rank| black_box(&digest).quantile(black_box(rank)))); +} + +#[divan::bench(args = [2, 6, 100])] +fn quantiles_batch(bencher: Bencher, num_ranks: usize) { + let digest = prepared_digest(); + let ranks = (1..=num_ranks) + .map(|rank| rank as f64 / (num_ranks + 1) as f64) + .collect::>(); + + bencher + .counter(ItemsCount::new(num_ranks)) + .bench_local(|| black_box(&digest).quantiles(black_box(&ranks))); +} diff --git a/datasketches/src/tdigest/sketch.rs b/datasketches/src/tdigest/sketch.rs index d496a9c5..80e84c55 100644 --- a/datasketches/src/tdigest/sketch.rs +++ b/datasketches/src/tdigest/sketch.rs @@ -506,6 +506,21 @@ impl TDigestMut { self.view().quantile(rank) } + /// Returns the quantiles described by [`TDigest::quantiles`]. + /// + /// # Panics + /// + /// Panics if any rank is outside `[0.0, 1.0]`. + pub fn quantiles(&mut self, ranks: &[f64]) -> Option> { + check_ranks(ranks); + + if self.is_empty() { + return None; + } + + self.view().quantiles(ranks) + } + /// Serializes this mutable t-digest to bytes. /// /// # Examples @@ -1268,6 +1283,35 @@ impl TDigest { self.view().quantile(rank) } + /// Computes approximate quantiles for the given normalized ranks. + /// + /// Ranks in nondecreasing order are answered with one centroid scan. Ranks in any other order + /// are accepted and results are returned in the same order as the input. + /// + /// Returns `None` if this t-digest is empty. + /// + /// # Panics + /// + /// Panics if any rank is outside `[0.0, 1.0]`. + /// + /// # Examples + /// + /// ``` + /// use datasketches::tdigest::TDigestMut; + /// + /// let mut sketch = TDigestMut::new(100).unwrap(); + /// for value in [1.0, 2.0, 3.0] { + /// sketch.update(value); + /// } + /// let digest = sketch.freeze(); + /// let quantiles = digest.quantiles(&[0.25, 0.5, 0.75]).unwrap(); + /// assert_eq!(quantiles.len(), 3); + /// ``` + pub fn quantiles(&self, ranks: &[f64]) -> Option> { + check_ranks(ranks); + self.view().quantiles(ranks) + } + /// Converts this immutable t-digest into a mutable one. /// /// # Examples @@ -1435,82 +1479,135 @@ impl TDigestView<'_> { return None; } + Some(QuantileCursor::new(self).quantile(rank)) + } + + fn quantiles(&self, ranks: &[f64]) -> Option> { + debug_assert!( + ranks.iter().all(|rank| (0.0..=1.0).contains(rank)), + "ranks must be in [0.0, 1.0]" + ); + + if self.centroids.is_empty() { + return None; + } + + let mut quantiles = vec![0.; ranks.len()]; + let mut cursor = QuantileCursor::new(self); + if ranks.windows(2).all(|pair| pair[0] <= pair[1]) { + for (index, &rank) in ranks.iter().enumerate() { + quantiles[index] = cursor.quantile(rank); + } + return Some(quantiles); + } + + // The cursor only moves forward. Sort indices so queries become monotonic without changing + // the caller's output order. + let mut rank_order = (0..ranks.len()).collect::>(); + rank_order.sort_by(|&left, &right| ranks[left].total_cmp(&ranks[right])); + for index in rank_order { + quantiles[index] = cursor.quantile(ranks[index]); + } + Some(quantiles) + } +} + +/// Incrementally answers quantile queries supplied in nondecreasing rank order. +struct QuantileCursor<'a> { + min: f64, + max: f64, + centroids: &'a [Centroid], + centroids_weight: f64, + centroid_index: usize, + weight_so_far: f64, +} + +impl<'a> QuantileCursor<'a> { + fn new(view: &TDigestView<'a>) -> Self { + debug_assert!(!view.centroids.is_empty()); + + QuantileCursor { + min: view.min, + max: view.max, + centroids: view.centroids, + centroids_weight: view.centroids_weight as f64, + centroid_index: 0, + weight_so_far: view.centroids[0].weight() / 2., + } + } + + fn quantile(&mut self, rank: f64) -> f64 { + debug_assert!(!self.centroids.is_empty()); + if self.centroids.len() == 1 { - return Some(self.centroids[0].mean); + return self.centroids[0].mean; } // at least 2 centroids - let centroids_weight = self.centroids_weight as f64; + let centroids_weight = self.centroids_weight; let num_centroids = self.centroids.len(); let weight = rank * centroids_weight; if weight < 1. { - return Some(self.min); + return self.min; } if weight > centroids_weight - 1. { - return Some(self.max); + return self.max; } let first_weight = self.centroids[0].weight(); if first_weight > 1. && weight < first_weight / 2. { - return Some( - self.min - + (((weight - 1.) / ((first_weight / 2.) - 1.)) - * (self.centroids[0].mean - self.min)), - ); + return self.min + + (((weight - 1.) / ((first_weight / 2.) - 1.)) + * (self.centroids[0].mean - self.min)); } let last_weight = self.centroids[num_centroids - 1].weight(); if last_weight > 1. && (centroids_weight - weight <= last_weight / 2.) { if last_weight == 2. { - return Some(self.max); + return self.max; } - return Some( - self.max - - (((centroids_weight - weight - 1.) / ((last_weight / 2.) - 1.)) - * (self.max - self.centroids[num_centroids - 1].mean)), - ); + return self.max + - (((centroids_weight - weight - 1.) / ((last_weight / 2.) - 1.)) + * (self.max - self.centroids[num_centroids - 1].mean)); } // interpolate between extremes - let mut weight_so_far = first_weight / 2.; - for i in 0..(num_centroids - 1) { - let dw = (self.centroids[i].weight() + self.centroids[i + 1].weight()) / 2.; - if weight_so_far + dw > weight { + while self.centroid_index < num_centroids - 1 { + let dw = (self.centroids[self.centroid_index].weight() + + self.centroids[self.centroid_index + 1].weight()) + / 2.; + if self.weight_so_far + dw > weight { // the target weight is between centroids i and i+1 let mut left_weight = 0.; - if self.centroids[i].weight.get() == 1 { - if weight - weight_so_far < 0.5 { - return Some(self.centroids[i].mean); + if self.centroids[self.centroid_index].weight.get() == 1 { + if weight - self.weight_so_far < 0.5 { + return self.centroids[self.centroid_index].mean; } left_weight = 0.5; } let mut right_weight = 0.; - if self.centroids[i + 1].weight.get() == 1 { - if weight_so_far + dw - weight <= 0.5 { - return Some(self.centroids[i + 1].mean); + if self.centroids[self.centroid_index + 1].weight.get() == 1 { + if self.weight_so_far + dw - weight <= 0.5 { + return self.centroids[self.centroid_index + 1].mean; } right_weight = 0.5; } // Each centroid is weighted by the distance from the target to the *other* // centroid, so the estimate approaches the nearer one. - let distance_from_left = weight - weight_so_far - left_weight; - let distance_to_right = weight_so_far + dw - weight - right_weight; - return Some(weighted_average( - self.centroids[i].mean, + let distance_from_left = weight - self.weight_so_far - left_weight; + let distance_to_right = self.weight_so_far + dw - weight - right_weight; + return weighted_average( + self.centroids[self.centroid_index].mean, distance_to_right, - self.centroids[i + 1].mean, + self.centroids[self.centroid_index + 1].mean, distance_from_left, - )); + ); } - weight_so_far += dw; + self.weight_so_far += dw; + self.centroid_index += 1; } let w1 = weight - (centroids_weight) - ((self.centroids[num_centroids - 1].weight()) / 2.); let w2 = (self.centroids[num_centroids - 1].weight() / 2.) - w1; - Some(weighted_average( - self.centroids[num_centroids - 1].mean, - w1, - self.max, - w2, - )) + weighted_average(self.centroids[num_centroids - 1].mean, w1, self.max, w2) } } @@ -1526,6 +1623,14 @@ fn check_split_points(split_points: &[f64]) { } } +#[track_caller] +fn check_ranks(ranks: &[f64]) { + assert!( + ranks.iter().all(|rank| (0.0..=1.0).contains(rank)), + "ranks must be in [0.0, 1.0]" + ); +} + fn centroid_cmp(a: &Centroid, b: &Centroid) -> Ordering { match a.mean.partial_cmp(&b.mean) { Some(order) => order, diff --git a/tests-integration/tests/tdigest_test/sketch.rs b/tests-integration/tests/tdigest_test/sketch.rs index 302cc4c7..efdc6436 100644 --- a/tests-integration/tests/tdigest_test/sketch.rs +++ b/tests-integration/tests/tdigest_test/sketch.rs @@ -393,6 +393,38 @@ fn test_quantile_handles_two_sample_last_centroid() { assert_eq!(tdigest.quantile(0.75), Some(100.0)); } +#[test] +fn test_batch_quantiles_match_scalar_queries_in_input_order() { + let mut tdigest = TDigestMut::new(100).unwrap(); + for value in 0..10_000 { + tdigest.update(((value * 37) % 1_003) as f64); + } + + for ranks in [ + vec![0.0, 0.001, 0.25, 0.5, 0.5, 0.99, 1.0], + vec![0.99, 0.0, 0.5, 1.0, 0.001, 0.5, 0.25], + vec![], + ] { + let expected = ranks + .iter() + .map(|&rank| tdigest.quantile(rank).unwrap()) + .collect::>(); + assert_eq!(tdigest.quantiles(&ranks), Some(expected.clone())); + assert_eq!(tdigest.clone().freeze().quantiles(&ranks), Some(expected)); + } +} + +#[test] +fn test_batch_quantiles_reject_invalid_ranks() { + let mut tdigest = TDigestMut::default(); + tdigest.update(1.0); + let tdigest = tdigest.freeze(); + + for ranks in [[-f64::EPSILON], [1.0 + f64::EPSILON], [f64::NAN]] { + assert!(std::panic::catch_unwind(|| tdigest.quantiles(&ranks)).is_err()); + } +} + #[test] fn test_rank_left_tail_is_a_fraction_of_the_total_weight() { let mut tdigest =